use std::ffi::OsString;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf, ReadHalf, WriteHalf};
use tokio::net::windows::named_pipe::{NamedPipeClient, NamedPipeServer};
use weida_core::{Error, PeerIdentity, WindowsPrincipal};
use weida_runtime::Exec;
use crate::chunked::{ChunkReader, ChunkWriter, Marker};
use crate::grouped::Stream;
pub(crate) struct PipeEndpoint {
pub(crate) path: OsString,
pub(crate) exec: Exec,
}
pub(crate) struct PipeStream {
io: PipeIo,
exec: Exec,
}
pub(crate) enum PipeIo {
Server(NamedPipeServer),
Client(NamedPipeClient),
}
impl PipeStream {
pub(crate) fn accepted(server: NamedPipeServer, exec: Exec) -> PipeStream {
PipeStream {
io: PipeIo::Server(server),
exec,
}
}
}
impl AsyncRead for PipeIo {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
match self.get_mut() {
PipeIo::Server(io) => Pin::new(io).poll_read(cx, buf),
PipeIo::Client(io) => Pin::new(io).poll_read(cx, buf),
}
}
}
impl AsyncWrite for PipeIo {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
match self.get_mut() {
PipeIo::Server(io) => Pin::new(io).poll_write(cx, buf),
PipeIo::Client(io) => Pin::new(io).poll_write(cx, buf),
}
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
PipeIo::Server(io) => Pin::new(io).poll_flush(cx),
PipeIo::Client(io) => Pin::new(io).poll_flush(cx),
}
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
PipeIo::Server(io) => Pin::new(io).poll_shutdown(cx),
PipeIo::Client(io) => Pin::new(io).poll_shutdown(cx),
}
}
}
impl AsyncRead for PipeStream {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().io).poll_read(cx, buf)
}
}
impl AsyncWrite for PipeStream {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.get_mut().io).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().io).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.get_mut().io).poll_shutdown(cx)
}
}
impl Stream for PipeStream {
type Endpoint = PipeEndpoint;
type Principal = WindowsPrincipal;
type Writer = ChunkWriter<WriteHalf<PipeIo>>;
type Reader = ChunkReader<ReadHalf<PipeIo>>;
async fn connect(endpoint: &PipeEndpoint) -> Result<PipeStream, Error> {
let client = weida_runtime::connect_pipe(&endpoint.exec, &endpoint.path).await?;
Ok(PipeStream {
io: PipeIo::Client(client),
exec: endpoint.exec.clone(),
})
}
fn principal(&self) -> Result<WindowsPrincipal, Error> {
match &self.io {
PipeIo::Server(io) => weida_runtime::client_principal(io),
PipeIo::Client(io) => weida_runtime::server_principal(io),
}
}
fn split(self) -> (Self::Reader, Self::Writer) {
let (recv, send) = tokio::io::split(self.io);
(
ChunkReader::new(recv, self.exec.clone()),
ChunkWriter::new(send, self.exec),
)
}
fn same_peer(group: &WindowsPrincipal, asking: &WindowsPrincipal) -> bool {
if group.sid != asking.sid {
return false;
}
match (group.pid, asking.pid) {
(Some(expected), Some(actual)) => expected == actual,
_ => true,
}
}
fn identity(principal: &WindowsPrincipal) -> PeerIdentity {
PeerIdentity::Windows(principal.clone())
}
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)
}
}