use std::fmt;
use std::io;
use std::os::windows::io::AsRawHandle;
use std::pin::Pin;
use std::ptr;
#[cfg(test)]
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll};
use ::tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use ::tokio::net::windows::named_pipe::NamedPipeServer;
use windows_sys::Win32::System::IO::CancelIoEx;
#[cfg(test)]
use crate::backend::BackendKind;
use crate::core::is_disconnect_error;
use crate::core::pseudocon::ConsoleShared;
use crate::core::session::Session as SessionCore;
#[cfg(test)]
use crate::error::Result;
#[cfg(test)]
use crate::size::Size;
use super::builder::PtyBuilder;
use crate::PtyController;
pub(crate) struct Pty {
pub(super) reader: ConoutReader,
pub(super) writer: ConinWriter,
pub(super) inner: Arc<SessionCore>,
}
impl fmt::Debug for Pty {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Pty")
.field("size", &self.inner.size())
.field("backend_kind", self.inner.backend_kind())
.finish_non_exhaustive()
}
}
impl Pty {
#[must_use]
pub(crate) fn builder() -> PtyBuilder {
PtyBuilder::default()
}
#[cfg(test)]
pub(crate) fn resize(&self, size: Size) -> Result<()> {
self.inner.resize(size)
}
#[must_use]
#[cfg(test)]
pub(crate) fn size(&self) -> Size {
self.inner.size()
}
#[cfg(test)]
pub(crate) fn clear(&self) -> Result<()> {
self.inner.clear()
}
#[must_use]
#[cfg(test)]
pub(crate) fn supports_clear(&self) -> bool {
self.inner.supports_clear()
}
#[must_use]
#[cfg(test)]
pub(crate) fn supports_release(&self) -> bool {
self.inner.supports_release()
}
#[must_use]
#[cfg(test)]
pub(crate) fn backend_kind(&self) -> &BackendKind {
self.inner.backend_kind()
}
#[must_use]
pub(crate) fn controller(&self) -> PtyController {
PtyController::new(Arc::clone(&self.inner))
}
#[must_use]
#[cfg(test)]
pub(crate) fn split(&mut self) -> (ReadHalf<'_>, WriteHalf<'_>) {
let Self { reader, writer, .. } = self;
(ReadHalf { reader }, WriteHalf { writer })
}
#[must_use]
pub(crate) fn into_split(self) -> (OwnedReadHalf, OwnedWriteHalf) {
let Self {
reader,
writer,
inner,
} = self;
let read_session = Arc::clone(&inner);
(
OwnedReadHalf {
reader,
_session: read_session,
},
OwnedWriteHalf {
writer,
_session: inner,
},
)
}
}
impl AsyncRead for Pty {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().reader).poll_read(cx, buf)
}
}
impl AsyncWrite for Pty {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.get_mut().writer).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().writer).poll_shutdown(cx)
}
}
#[cfg(test)]
pub(crate) struct ReadHalf<'a> {
reader: &'a mut ConoutReader,
}
#[cfg(test)]
impl fmt::Debug for ReadHalf<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ReadHalf").finish_non_exhaustive()
}
}
#[cfg(test)]
impl AsyncRead for ReadHalf<'_> {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Pin::new(&mut *self.get_mut().reader).poll_read(cx, buf)
}
}
#[cfg(test)]
pub(crate) struct WriteHalf<'a> {
writer: &'a mut ConinWriter,
}
#[cfg(test)]
impl fmt::Debug for WriteHalf<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WriteHalf").finish_non_exhaustive()
}
}
#[cfg(test)]
impl AsyncWrite for WriteHalf<'_> {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut *self.get_mut().writer).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut *self.get_mut().writer).poll_shutdown(cx)
}
}
pub struct OwnedReadHalf {
reader: ConoutReader,
_session: Arc<SessionCore>,
}
impl fmt::Debug for OwnedReadHalf {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OwnedReadHalf").finish_non_exhaustive()
}
}
impl AsyncRead for OwnedReadHalf {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().reader).poll_read(cx, buf)
}
}
pub struct OwnedWriteHalf {
writer: ConinWriter,
_session: Arc<SessionCore>,
}
impl fmt::Debug for OwnedWriteHalf {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OwnedWriteHalf").finish_non_exhaustive()
}
}
impl AsyncWrite for OwnedWriteHalf {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.get_mut().writer).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.get_mut().writer).poll_shutdown(cx)
}
}
#[derive(Debug)]
pub(super) struct ConoutReader {
pipe: Option<NamedPipeServer>,
shared: Arc<ConsoleShared>,
saw_eof: bool,
}
fn notify_eof_once(saw_eof: &mut bool, notify: impl FnOnce()) {
if !*saw_eof {
*saw_eof = true;
notify();
}
}
fn conout_error_as_eof(err: io::Error) -> io::Result<()> {
if is_disconnect_error(&err) {
Ok(())
} else {
Err(err)
}
}
impl ConoutReader {
pub(super) const fn new(pipe: NamedPipeServer, shared: Arc<ConsoleShared>) -> Self {
Self {
pipe: Some(pipe),
shared,
saw_eof: false,
}
}
}
impl AsyncRead for ConoutReader {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
if buf.remaining() == 0 {
return Poll::Ready(Ok(()));
}
let this = self.get_mut();
let Some(pipe) = this.pipe.as_mut() else {
return Poll::Ready(Ok(()));
};
let before = buf.filled().len();
match Pin::new(pipe).poll_read(cx, buf) {
Poll::Ready(Ok(())) => {
if buf.filled().len() == before {
let shared = &this.shared;
notify_eof_once(&mut this.saw_eof, || shared.notify_reader_eof());
}
Poll::Ready(Ok(()))
},
Poll::Ready(Err(err)) => match conout_error_as_eof(err) {
Ok(()) => {
let shared = &this.shared;
notify_eof_once(&mut this.saw_eof, || shared.notify_reader_eof());
Poll::Ready(Ok(()))
},
Err(err) => Poll::Ready(Err(err)),
},
Poll::Pending => Poll::Pending,
}
}
}
impl Drop for ConoutReader {
fn drop(&mut self) {
drop(self.pipe.take());
self.shared.notify_reader_closed();
}
}
#[derive(Debug)]
pub(super) struct ConinWriter {
pipe: Option<NamedPipeServer>,
session: Arc<SessionCore>,
#[cfg(test)]
close_observer: Option<Arc<AtomicBool>>,
}
impl ConinWriter {
pub(super) const fn new(pipe: NamedPipeServer, session: Arc<SessionCore>) -> Self {
Self {
pipe: Some(pipe),
session,
#[cfg(test)]
close_observer: None,
}
}
fn pipe(&mut self) -> io::Result<&mut NamedPipeServer> {
self.pipe.as_mut().ok_or_else(|| {
io::Error::new(
io::ErrorKind::BrokenPipe,
"the pseudoconsole input pipe has been shut down",
)
})
}
fn close_pipe(&mut self) {
let Some(pipe) = self.pipe.take() else {
return;
};
#[cfg(test)]
if let Some(observer) = &self.close_observer {
observer.store(true, Ordering::SeqCst);
}
unsafe { CancelIoEx(pipe.as_raw_handle(), ptr::null()) };
drop(pipe);
self.session.request_close_after_input();
}
}
impl Drop for ConinWriter {
fn drop(&mut self) {
self.close_pipe();
}
}
impl AsyncWrite for ConinWriter {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let pipe = match self.get_mut().pipe() {
Ok(pipe) => pipe,
Err(err) => return Poll::Ready(Err(err)),
};
Pin::new(pipe).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
self.get_mut().close_pipe();
Poll::Ready(Ok(()))
}
}
#[cfg(test)]
mod behavior_tests {
use std::cell::Cell;
use std::io;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use super::{conout_error_as_eof, notify_eof_once, Pty};
#[test]
fn eof_notification_runs_exactly_once() {
let mut saw_eof = false;
let notifications = Cell::new(0);
notify_eof_once(&mut saw_eof, || notifications.set(notifications.get() + 1));
notify_eof_once(&mut saw_eof, || notifications.set(notifications.get() + 1));
assert!(saw_eof);
assert_eq!(notifications.get(), 1);
}
#[test]
fn only_disconnect_errors_become_eof() {
assert!(conout_error_as_eof(io::Error::new(io::ErrorKind::BrokenPipe, "closed")).is_ok());
let err = conout_error_as_eof(io::Error::new(io::ErrorKind::PermissionDenied, "denied"))
.expect_err("an unrelated I/O failure must not become EOF");
assert_eq!(err.kind(), io::ErrorKind::PermissionDenied);
}
#[tokio::test]
async fn dropping_a_pty_runs_the_conin_cancellation_path() {
let mut pty = Pty::builder().build().expect("building must succeed");
let closed = Arc::new(AtomicBool::new(false));
pty.writer.close_observer = Some(Arc::clone(&closed));
drop(pty);
assert!(
closed.load(Ordering::SeqCst),
"ConinWriter::drop must cancel pending I/O before closing the pipe"
);
}
}