#[cfg(unix)]
mod imp {
use std::os::fd::AsRawFd;
use std::path::PathBuf;
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{UnixListener, UnixStream};
use tokio::sync::watch;
use tokio_util::bytes::{Bytes, BytesMut};
use tokio_util::codec::{Decoder, Encoder, LengthDelimitedCodec};
use crate::comms::daemon::Broker;
use crate::comms::protocol::{CommsOut, CommsRequest};
use crate::comms::transport::{CommsFrontend, CommsLink, MAX_FRAME_BYTES, PeerCred, serve_link};
const READ_CHUNK: usize = 8 * 1024;
pub struct UdsLink {
stream: UnixStream,
codec: LengthDelimitedCodec,
read_buf: BytesMut,
peer: PeerCred,
}
impl UdsLink {
fn new(stream: UnixStream, peer: PeerCred) -> Self {
let mut codec = LengthDelimitedCodec::new();
codec.set_max_frame_length(MAX_FRAME_BYTES);
Self {
stream,
codec,
read_buf: BytesMut::with_capacity(READ_CHUNK),
peer,
}
}
}
impl CommsLink for UdsLink {
async fn recv(&mut self) -> std::io::Result<Option<CommsRequest>> {
loop {
if let Some(frame) = self.codec.decode(&mut self.read_buf)? {
let req = rmp_serde::from_slice(&frame)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
return Ok(Some(req));
}
let n = self.stream.read_buf(&mut self.read_buf).await?;
if n == 0 {
if self.read_buf.is_empty() {
return Ok(None);
}
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"peer closed mid-frame",
));
}
}
}
async fn send(&mut self, out: CommsOut) -> std::io::Result<()> {
let body = rmp_serde::to_vec_named(&out)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e.to_string()))?;
let mut framed = BytesMut::new();
self.codec.encode(Bytes::from(body), &mut framed)?;
self.stream.write_all(&framed).await?;
self.stream.flush().await
}
fn peer_cred(&self) -> PeerCred {
self.peer
}
}
pub struct UdsFrontend {
listener: UnixListener,
socket_path: PathBuf,
}
impl UdsFrontend {
pub fn from_listener(listener: UnixListener, socket_path: PathBuf) -> Self {
Self { listener, socket_path }
}
}
impl CommsFrontend for UdsFrontend {
async fn serve(
self: Box<Self>,
broker: Arc<Broker>,
mut shutdown: watch::Receiver<bool>,
) -> std::io::Result<()> {
broker.mark_active().await;
let my_uid = super::daemon_uid();
loop {
tokio::select! {
accepted = self.listener.accept() => {
let (stream, _addr) = match accepted {
Ok(pair) => pair,
Err(e) => {
tracing::warn!(error = %e, "comms: accept failed");
continue;
}
};
let peer = peer_cred_of(&stream);
if let Some(uid) = peer.uid && uid != my_uid {
tracing::warn!(
peer_uid = uid,
daemon_uid = my_uid,
"comms: rejecting cross-user connection"
);
continue;
}
let guard = broker.register_link();
let broker = broker.clone();
tokio::spawn(async move {
let is_relay =
peek_first_byte(&stream).await == Some(crate::comms::relay::RELAY_MAGIC[0]);
if is_relay {
broker.serve_relay_connection(stream, guard).await;
} else {
serve_link(broker, UdsLink::new(stream, peer), guard).await;
}
});
}
_ = shutdown.changed() => {
if *shutdown.borrow() {
break;
}
}
}
}
let _ = std::fs::remove_file(&self.socket_path);
broker.drain_links(crate::comms::daemon::DRAIN_GRACE).await;
Ok(())
}
}
fn peer_cred_of(stream: &UnixStream) -> PeerCred {
super::peer_cred_from_fd(stream.as_raw_fd())
}
async fn peek_first_byte(stream: &UnixStream) -> Option<u8> {
use tokio::io::Interest;
loop {
stream.readable().await.ok()?;
let mut byte = 0u8;
let peeked = stream.try_io(Interest::READABLE, || {
let n = unsafe {
super::recv(
stream.as_raw_fd(),
std::ptr::from_mut(&mut byte).cast(),
1,
super::MSG_PEEK,
)
};
if n < 0 {
Err(std::io::Error::last_os_error())
} else {
Ok(n)
}
});
match peeked {
Ok(0) => return None,
Ok(_) => return Some(byte),
Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => continue,
Err(_) => return None,
}
}
}
}
#[cfg(unix)]
pub use imp::{UdsFrontend, UdsLink};
#[cfg(unix)]
pub fn daemon_uid() -> u32 {
unsafe { getuid() }
}
#[cfg(not(unix))]
pub fn daemon_uid() -> u32 {
0
}
#[cfg(unix)]
const MSG_PEEK: i32 = 0x2;
#[cfg(unix)]
unsafe extern "C" {
fn getuid() -> u32;
fn getsockopt(sockfd: i32, level: i32, optname: i32, optval: *mut core::ffi::c_void, optlen: *mut u32) -> i32;
fn recv(sockfd: i32, buf: *mut core::ffi::c_void, len: usize, flags: i32) -> isize;
}
#[cfg(unix)]
pub(crate) fn peer_cred_from_fd(fd: i32) -> crate::comms::transport::PeerCred {
#[cfg(target_os = "linux")]
{
const SOL_SOCKET: i32 = 1;
const SO_PEERCRED: i32 = 17;
#[repr(C)]
#[derive(Default, Clone, Copy)]
struct Ucred {
pid: i32,
uid: u32,
gid: u32,
}
let mut cred = Ucred::default();
let mut len = core::mem::size_of::<Ucred>() as u32;
let rc = unsafe { getsockopt(fd, SOL_SOCKET, SO_PEERCRED, (&mut cred as *mut Ucred).cast(), &mut len) };
if rc == 0 {
return crate::comms::transport::PeerCred {
uid: Some(cred.uid),
pid: Some(cred.pid as u32),
};
}
}
#[cfg(target_os = "macos")]
{
const SOL_LOCAL: i32 = 0;
const LOCAL_PEERCRED: i32 = 0x001;
#[repr(C)]
struct Xucred {
cr_version: u32,
cr_uid: u32,
cr_ngroups: i16,
cr_groups: [u32; 16],
}
let mut cred = Xucred {
cr_version: 0,
cr_uid: u32::MAX,
cr_ngroups: 0,
cr_groups: [0; 16],
};
let mut len = core::mem::size_of::<Xucred>() as u32;
let rc = unsafe {
getsockopt(
fd,
SOL_LOCAL,
LOCAL_PEERCRED,
(&mut cred as *mut Xucred).cast(),
&mut len,
)
};
if rc == 0 {
return crate::comms::transport::PeerCred {
uid: Some(cred.cr_uid),
pid: None,
};
}
}
crate::comms::transport::PeerCred::default()
}
#[cfg(all(test, unix))]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{UnixListener, UnixStream};
use tokio::sync::watch;
use super::UdsFrontend;
use crate::comms::daemon::Broker;
use crate::comms::protocol::{CommsOut, CommsRequest, CommsResponse};
use crate::comms::store::CommsStore;
use crate::comms::transport::CommsFrontend;
async fn send_req(stream: &mut UnixStream, req: &CommsRequest) {
let body = rmp_serde::to_vec_named(req).expect("encode");
let len = u32::try_from(body.len()).expect("len fits");
stream.write_all(&len.to_be_bytes()).await.expect("write len");
stream.write_all(&body).await.expect("write body");
stream.flush().await.expect("flush");
}
async fn read_resp(stream: &mut UnixStream) -> CommsOut {
let mut prefix = [0u8; 4];
stream.read_exact(&mut prefix).await.expect("read len");
let len = u32::from_be_bytes(prefix) as usize;
let mut buf = vec![0u8; len];
stream.read_exact(&mut buf).await.expect("read body");
rmp_serde::from_slice(&buf).expect("decode")
}
#[tokio::test]
async fn a_silent_client_does_not_stall_the_accept_loop() {
let dir = tempfile::tempdir().expect("tempdir");
let store = Arc::new(CommsStore::open(dir.path()).expect("store"));
let broker = Arc::new(Broker::new(store));
let socket = dir.path().join("accept.sock");
let listener = UnixListener::bind(&socket).expect("bind");
let frontend = UdsFrontend::from_listener(listener, socket.clone());
let (_shutdown_tx, shutdown_rx) = watch::channel(false);
tokio::spawn(Box::new(frontend).serve(broker, shutdown_rx));
let _silent = UnixStream::connect(&socket).await.expect("connect A");
tokio::time::sleep(Duration::from_millis(200)).await;
let mut b = UnixStream::connect(&socket).await.expect("connect B");
send_req(&mut b, &CommsRequest::Ping).await;
let resp = tokio::time::timeout(Duration::from_secs(3), read_resp(&mut b))
.await
.expect("B must be served even while A is silent — otherwise the accept loop is stalled");
assert_eq!(
resp,
CommsOut::Response(CommsResponse::Pong),
"B's Ping must be answered with Pong"
);
}
}