#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod tests;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::thread;
use ohno::AppError;
use crate::constants::CONNECT_TIMEOUT;
use crate::output::note_line;
use crate::pal::error::{PalError, PalErrorKind};
use crate::pal::ids::{ConnId, RelayLeaseId};
use crate::pal::local_console::{ConsoleInput, LocalConsole};
use crate::pal::transport::Transport;
use crate::protocol::Message;
use crate::{
AttachFailedError, ConsoleRestoreError, DisplacedError, NoConsoleError, Outcome,
PalFailedError, RelayFailedError, ResumeTimeoutError, SessionId, SupervisorLostError,
};
pub(crate) fn attach<T, C>(
transport: &T,
console: &C,
pipe_name: &str,
session_id: SessionId,
) -> Result<Outcome, AppError>
where
T: Transport + Clone + Send + Sync + 'static,
C: LocalConsole + Clone + Send + Sync + 'static,
{
if !console.has_console() {
return Err(NoConsoleError::new().into());
}
let lease = ConsoleLease::take(console).map_err(PalFailedError::caused_by)?;
let outcome = handshake_and_relay(transport, console, pipe_name, session_id);
let restored = lease.release();
finish_with_cleanup(outcome, restored)
}
fn finish_with_cleanup(
outcome: Result<Outcome, AppError>,
cleaned_up: Result<(), AppError>,
) -> Result<Outcome, AppError> {
match (outcome, cleaned_up) {
(outcome, Ok(())) => outcome,
(Ok(Outcome::AppExit(status)), Err(error)) => {
note_line(format_args!("Warning: {error}"));
Ok(Outcome::AppExit(status))
}
(Ok(Outcome::Success), Err(error)) => Err(error),
(Err(error), Err(_restore_error)) => Err(error),
}
}
fn handshake_and_relay<T, C>(
transport: &T,
console: &C,
pipe_name: &str,
session_id: SessionId,
) -> Result<Outcome, AppError>
where
T: Transport + Clone + Send + Sync + 'static,
C: LocalConsole + Clone + Send + Sync + 'static,
{
let size = console.window_size().map_err(PalFailedError::caused_by)?;
let conn = match transport.connect(pipe_name, CONNECT_TIMEOUT) {
Ok(conn) => conn,
Err(error) if error.kind() == PalErrorKind::Timeout => {
return Err(ResumeTimeoutError::for_id(session_id).into());
}
Err(_) => return Err(AttachFailedError::for_id(session_id).into()),
};
transport
.send(conn, &Message::Attach { size })
.map_err(|_error| AttachFailedError::for_id(session_id))?;
match transport.recv(conn) {
Ok(Message::Attached {
session_id: attached_id,
}) if attached_id == session_id => {}
Ok(Message::Displaced) => {
transport.disconnect(conn);
return Err(DisplacedError::new().into());
}
_ => {
transport.disconnect(conn);
return Err(AttachFailedError::for_id(session_id).into());
}
}
relay(transport, console, conn)
}
struct ConsoleLease<'a, C: LocalConsole> {
console: &'a C,
lease: Option<RelayLeaseId>,
}
impl<'a, C: LocalConsole> ConsoleLease<'a, C> {
fn take(console: &'a C) -> Result<Self, PalError> {
let lease = console.begin_raw_relay()?;
Ok(Self {
console,
lease: Some(lease),
})
}
#[cfg_attr(test, mutants::skip)]
fn release(mut self) -> Result<(), AppError> {
self.lease.take().map_or(Ok(()), |lease| {
self.console
.end_raw_relay(lease)
.map_err(|error| ConsoleRestoreError::caused_by(error).into())
})
}
}
impl<C: LocalConsole> Drop for ConsoleLease<'_, C> {
fn drop(&mut self) {
if let Some(lease) = self.lease.take() {
_ = self.console.end_raw_relay(lease);
}
}
}
fn relay<T, C>(transport: &T, console: &C, conn: ConnId) -> Result<Outcome, AppError>
where
T: Transport + Clone + Send + Sync + 'static,
C: LocalConsole + Clone + Send + Sync + 'static,
{
let input_failed = Arc::new(AtomicBool::new(false));
let reader = spawn_input_reader(transport, console, conn, &input_failed);
let outcome = receive_until_relay_ends(transport, console, conn);
let input_failed = input_failed.load(Ordering::SeqCst);
let cancelled = console
.cancel_input()
.map_err(PalFailedError::caused_by)
.map_err(AppError::from);
if cancelled.is_ok() {
_ = reader.join();
}
let outcome = match outcome {
RelayEnd::AppExited(status) => Ok(Outcome::AppExit(status)),
RelayEnd::Displaced => Err(DisplacedError::new().into()),
RelayEnd::Failed => Err(RelayFailedError::new().into()),
RelayEnd::SupervisorGone => {
if input_failed {
Err(RelayFailedError::new().into())
} else {
Err(SupervisorLostError::new().into())
}
}
};
finish_with_cleanup(outcome, cancelled)
}
fn spawn_input_reader<T, C>(
transport: &T,
console: &C,
conn: ConnId,
input_failed: &Arc<AtomicBool>,
) -> thread::JoinHandle<()>
where
T: Transport + Clone + Send + Sync + 'static,
C: LocalConsole + Clone + Send + Sync + 'static,
{
thread::spawn({
let transport = transport.clone();
let console = console.clone();
let input_failed = Arc::clone(input_failed);
move || {
loop {
match console.read_input() {
Ok(ConsoleInput::Bytes(bytes)) => {
if transport.send(conn, &Message::Input(bytes)).is_err() {
break;
}
}
Ok(ConsoleInput::Resize(size)) => {
if transport.send(conn, &Message::Resize { size }).is_err() {
break;
}
}
Err(_) => {
input_failed.store(true, Ordering::SeqCst);
transport.disconnect(conn);
break;
}
}
}
}
})
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum RelayEnd {
AppExited(i32),
Displaced,
Failed,
SupervisorGone,
}
#[cfg_attr(test, mutants::skip)]
fn receive_until_relay_ends<T, C>(transport: &T, console: &C, conn: ConnId) -> RelayEnd
where
T: Transport + Clone + Send + Sync + 'static,
C: LocalConsole + Clone + Send + Sync + 'static,
{
let end = loop {
match transport.recv(conn) {
Ok(Message::Output(bytes)) => {
if console.write_output(&bytes).is_err() {
break RelayEnd::Failed;
}
}
Ok(Message::AppExited { status }) => break RelayEnd::AppExited(status),
Ok(Message::Displaced) => break RelayEnd::Displaced,
Err(error) if error.kind() == PalErrorKind::Disconnected => {
break RelayEnd::SupervisorGone;
}
Ok(_) | Err(_) => break RelayEnd::Failed,
}
};
transport.disconnect(conn);
end
}