use std::path::PathBuf;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::UnixStream;
use tokio::net::unix::{OwnedReadHalf, OwnedWriteHalf};
use weida_core::{Error, LocalPrincipal, LossCause, PeerIdentity};
use weida_runtime::{Exec, peer_credentials};
use crate::chunked::{ChunkReader, ChunkWriter, Marker};
use crate::grouped::Stream;
pub(crate) struct UnixEndpoint {
pub(crate) path: PathBuf,
pub(crate) exec: Exec,
}
pub(crate) struct UnixLocal {
io: UnixStream,
exec: Exec,
}
impl UnixLocal {
pub(crate) fn accepted(io: UnixStream, exec: Exec) -> UnixLocal {
UnixLocal { io, exec }
}
}
impl AsyncRead for UnixLocal {
fn poll_read(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.get_mut().io).poll_read(cx, buf)
}
}
impl AsyncWrite for UnixLocal {
fn poll_write(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
buf: &[u8],
) -> std::task::Poll<std::io::Result<usize>> {
std::pin::Pin::new(&mut self.get_mut().io).poll_write(cx, buf)
}
fn poll_flush(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.get_mut().io).poll_flush(cx)
}
fn poll_shutdown(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<std::io::Result<()>> {
std::pin::Pin::new(&mut self.get_mut().io).poll_shutdown(cx)
}
}
impl Stream for UnixLocal {
type Endpoint = UnixEndpoint;
type Principal = LocalPrincipal;
type Writer = ChunkWriter<OwnedWriteHalf>;
type Reader = ChunkReader<OwnedReadHalf>;
async fn connect(endpoint: &UnixEndpoint) -> Result<UnixLocal, Error> {
let io = UnixStream::connect(&endpoint.path)
.await
.map_err(|e| match e.kind() {
std::io::ErrorKind::NotFound | std::io::ErrorKind::ConnectionRefused => {
Error::ConnectionLost(LossCause::PeerClosed)
}
_ => Error::Io(e),
})?;
Ok(UnixLocal {
io,
exec: endpoint.exec.clone(),
})
}
fn principal(&self) -> Result<LocalPrincipal, Error> {
peer_credentials(&self.io)
}
fn split(self) -> (Self::Reader, Self::Writer) {
let (recv, send) = self.io.into_split();
(
ChunkReader::new(recv, self.exec.clone()),
ChunkWriter::new(send, self.exec),
)
}
fn same_peer(group: &LocalPrincipal, asking: &LocalPrincipal) -> bool {
if group.uid != asking.uid {
return false;
}
match (group.pid, asking.pid) {
(Some(expected), Some(actual)) => expected == actual,
_ => true,
}
}
fn identity(principal: &LocalPrincipal) -> PeerIdentity {
PeerIdentity::Local(*principal)
}
fn finish(writer: Self::Writer) {
writer.end(Marker::Fin);
}
fn reset(writer: Self::Writer, code: u64) {
writer.end(Marker::Reset(code));
}
fn stop(reader: Self::Reader, _code: u64) {
reader.drain();
}
fn read_error(error: std::io::Error) -> Error {
crate::chunked::read_error(error)
}
}