use std::collections::HashMap;
use std::pin::Pin;
use bytes::Bytes;
use futures::Stream;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::sync::{mpsc, Mutex};
use tokio::task::JoinHandle;
use tracing::{debug, warn};
use alkcall::channels::client::ChannelClient;
use alkcall::core::Connection;
use crate::control::ControlMessage;
use crate::negotiation::{NegotiateRequest, NegotiationError, NegotiationWriter};
use crate::wire::{Chunk, ChunkReader, ChunkWriter, RawError, STREAM_CTRL_IN, STREAM_CTRL_OUT};
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum TtySessionError {
#[error("io: {0}")]
Io(#[from] std::io::Error),
#[error("wire: {0}")]
Wire(#[from] RawError),
#[error("negotiation serialize: {0}")]
NegotiationSerialize(#[from] serde_json::Error),
#[error("negotiation write: {0}")]
NegotiationWrite(#[from] NegotiationError),
#[error("channels open: {0}")]
ChannelsOpen(#[from] alkcall::channels::client::ChannelOpenError),
#[error("invalid open params: {0}")]
InvalidParams(String),
#[error("negotiation rejected: {error}")]
NegotiationRejected {
error: String,
fields: HashMap<String, String>,
},
#[error("session ended without exit chunk")]
NoExitChunk,
#[error("malformed exit chunk: {0}")]
MalformedExitChunk(String),
#[error("session open failed: {0}")]
Open(#[from] alkcall::core::StreamError),
}
pub struct TtySession {
writer: Mutex<ChunkWriter<Box<dyn AsyncWrite + Send + Unpin>>>,
stdout_rx: Mutex<Option<mpsc::Receiver<Bytes>>>,
stderr_rx: Mutex<Option<mpsc::Receiver<Bytes>>>,
exit_code: tokio::sync::watch::Receiver<Option<ExitOutcome>>,
_read_pump: JoinHandle<()>,
}
#[derive(Debug, Clone)]
enum ExitOutcome {
Exited(i32),
MalformedExit(String),
NoExitChunk,
}
impl TtySession {
pub async fn connect_direct(
connection: Connection,
negotiate: NegotiateRequest,
) -> Result<Self, TtySessionError> {
let stream = connection.accept_bi().await?;
Self::from_bidi_stream(stream, negotiate).await
}
pub async fn open_via_channels(
client: &ChannelClient,
params: serde_json::Value,
) -> Result<Self, TtySessionError> {
use serde::Deserialize as _;
NegotiateRequest::deserialize(¶ms)
.map_err(|e| TtySessionError::InvalidParams(e.to_string()))?;
let (channel_id, send, recv) = client
.open_channel(
crate::channels::OP_TTY_OPEN,
params,
crate::channels::TTY_ALPN,
)
.await
.map_err(TtySessionError::from)?;
debug!("tty: opened channel {channel_id} via channels");
let remote_addr = client.manager().remote_addr();
let source = alkcall::channels::source::channel_source(recv, send, remote_addr);
let channel_conn =
Connection::from_source(source, crate::channels::TTY_ALPN.as_bytes().to_vec());
Self::from_bidi_stream_via(channel_conn).await
}
async fn from_bidi_stream(
stream: alkcall::core::BiStream,
negotiate: NegotiateRequest,
) -> Result<Self, TtySessionError> {
let (read, write) = tokio::io::split(stream);
Self::from_halves(read, write, negotiate).await
}
async fn from_bidi_stream_via(channel_conn: Connection) -> Result<Self, TtySessionError> {
let stream = channel_conn.accept_bi().await?;
let (read, write) = tokio::io::split(stream);
Self::from_halves_raw(read, write).await
}
async fn from_halves<R, W>(
read: R,
write: W,
negotiate: NegotiateRequest,
) -> Result<Self, TtySessionError>
where
R: AsyncRead + Send + Unpin + 'static,
W: AsyncWrite + Send + Unpin + 'static,
{
let boxed_write: Box<dyn AsyncWrite + Send + Unpin> = Box::new(write);
let mut neg_writer = NegotiationWriter::new(boxed_write);
let body = serde_json::to_vec(&negotiate)?;
neg_writer.write_frame(&body).await?;
let writer = ChunkWriter::new(neg_writer.into_inner());
let mut reader = ChunkReader::new(read);
match reader.peek_stream_type().await {
Ok(0x00) => {
return Err(read_negotiation_error(reader.into_inner()).await);
}
Ok(_) => {}
Err(RawError::ConnectionClosed) => {
}
Err(e) => return Err(TtySessionError::Wire(e)),
}
Self::start_pump(writer, reader)
}
async fn from_halves_raw<R, W>(read: R, write: W) -> Result<Self, TtySessionError>
where
R: AsyncRead + Send + Unpin + 'static,
W: AsyncWrite + Send + Unpin + 'static,
{
let writer = ChunkWriter::new(Box::new(write) as Box<dyn AsyncWrite + Send + Unpin>);
let mut reader = ChunkReader::new(read);
match reader.peek_stream_type().await {
Ok(0x00) => {
return Err(read_negotiation_error(reader.into_inner()).await);
}
Ok(_) => {}
Err(RawError::ConnectionClosed) => {
}
Err(e) => return Err(TtySessionError::Wire(e)),
}
Self::start_pump(writer, reader)
}
fn start_pump<R>(
writer: ChunkWriter<Box<dyn AsyncWrite + Send + Unpin>>,
reader: ChunkReader<R>,
) -> Result<Self, TtySessionError>
where
R: AsyncRead + Send + Unpin + 'static,
{
let (stdout_tx, stdout_rx) = mpsc::channel::<Bytes>(64);
let (stderr_tx, stderr_rx) = mpsc::channel::<Bytes>(64);
let (exit_tx, exit_rx) = tokio::sync::watch::channel::<Option<ExitOutcome>>(None);
let read_pump = tokio::spawn(read_pump(reader, stdout_tx, stderr_tx, exit_tx));
Ok(Self {
writer: Mutex::new(writer),
stdout_rx: Mutex::new(Some(stdout_rx)),
stderr_rx: Mutex::new(Some(stderr_rx)),
exit_code: exit_rx,
_read_pump: read_pump,
})
}
pub async fn send_stdin(&self, bytes: Bytes) -> Result<(), TtySessionError> {
let mut writer = self.writer.lock().await;
let chunk = Chunk::stdin(bytes);
writer.write_chunk(&chunk).await?;
Ok(())
}
pub async fn close_stdin(&self) -> Result<(), TtySessionError> {
let mut writer = self.writer.lock().await;
let chunk = Chunk::stdin(Bytes::new());
writer.write_chunk(&chunk).await?;
Ok(())
}
pub async fn resize(
&self,
cols: u16,
rows: u16,
pixel_width: u16,
pixel_height: u16,
) -> Result<(), TtySessionError> {
let mut writer = self.writer.lock().await;
let msg = ControlMessage::Resize {
cols,
rows,
pixel_width,
pixel_height,
};
let json = msg.to_json()?;
let chunk = Chunk::ctrl_in(json);
writer.write_chunk(&chunk).await?;
Ok(())
}
pub async fn signal(&self, name: &str) -> Result<(), TtySessionError> {
let mut writer = self.writer.lock().await;
let msg = ControlMessage::Signal {
name: name.to_string(),
};
let json = msg.to_json()?;
let chunk = Chunk::ctrl_in(json);
writer.write_chunk(&chunk).await?;
Ok(())
}
pub async fn recv_stdout(&self) -> Pin<Box<dyn Stream<Item = Bytes> + Send>> {
let mut guard = self.stdout_rx.lock().await;
if let Some(rx) = guard.take() {
return Box::pin(futures::stream::unfold(rx, |mut rx| async move {
match rx.recv().await {
Some(bytes) if bytes.is_empty() => None,
Some(bytes) => Some((bytes, rx)),
None => None,
}
}));
}
Box::pin(futures::stream::empty())
}
pub async fn recv_stderr(&self) -> Option<Pin<Box<dyn Stream<Item = Bytes> + Send>>> {
let mut guard = self.stderr_rx.lock().await;
let rx = guard.take()?;
Some(Box::pin(futures::stream::unfold(rx, |mut rx| async move {
rx.recv().await.map(|bytes| (bytes, rx))
})))
}
pub async fn wait(&self) -> Result<i32, TtySessionError> {
let mut rx = self.exit_code.clone();
{
let borrow = rx.borrow();
if let Some(outcome) = borrow.as_ref() {
return outcome_to_result(outcome);
}
}
rx.changed()
.await
.map_err(|_| TtySessionError::NoExitChunk)?;
let borrow = rx.borrow();
match borrow.as_ref() {
Some(outcome) => outcome_to_result(outcome),
None => Err(TtySessionError::NoExitChunk),
}
}
}
impl Drop for TtySession {
fn drop(&mut self) {
self._read_pump.abort();
}
}
fn outcome_to_result(outcome: &ExitOutcome) -> Result<i32, TtySessionError> {
match outcome {
ExitOutcome::Exited(code) => Ok(*code),
ExitOutcome::MalformedExit(msg) => Err(TtySessionError::MalformedExitChunk(msg.clone())),
ExitOutcome::NoExitChunk => Err(TtySessionError::NoExitChunk),
}
}
async fn read_negotiation_error<R>(mut read: R) -> TtySessionError
where
R: AsyncRead + Unpin,
{
use tokio::io::AsyncReadExt;
let mut len_rest = [0u8; 3];
if let Err(e) = read.read_exact(&mut len_rest).await {
return TtySessionError::Wire(RawError::Io(e));
}
let length = u32::from_be_bytes([0x00, len_rest[0], len_rest[1], len_rest[2]]) as usize;
let mut body = vec![0u8; length];
if let Err(e) = read.read_exact(&mut body).await {
return TtySessionError::Wire(RawError::Io(e));
}
let value: serde_json::Value = match serde_json::from_slice(&body) {
Ok(v) => v,
Err(_) => {
return TtySessionError::NegotiationRejected {
error: String::from_utf8_lossy(&body).into_owned(),
fields: HashMap::new(),
};
}
};
let mut fields = HashMap::new();
let mut error = String::new();
if let Some(obj) = value.as_object() {
for (k, v) in obj {
if k == "error" {
if let Some(s) = v.as_str() {
error = s.to_string();
}
} else if let Some(s) = v.as_str() {
fields.insert(k.clone(), s.to_string());
}
}
}
TtySessionError::NegotiationRejected { error, fields }
}
async fn read_pump<R>(
mut reader: ChunkReader<R>,
stdout_tx: mpsc::Sender<Bytes>,
stderr_tx: mpsc::Sender<Bytes>,
exit_tx: tokio::sync::watch::Sender<Option<ExitOutcome>>,
) where
R: AsyncRead + Send + Unpin + 'static,
{
let mut exit_resolved = false;
loop {
let read = reader.read_chunk().await;
match read {
Ok(chunk) => match chunk.stream_type {
crate::wire::STREAM_STDOUT => {
if stdout_tx.send(chunk.bytes).await.is_err() {
debug!("tty: stdout receiver dropped, ending read pump");
break;
}
}
crate::wire::STREAM_STDERR => {
if stderr_tx.send(chunk.bytes).await.is_err() {
debug!("tty: stderr receiver dropped, ending read pump");
break;
}
}
STREAM_CTRL_OUT => match ControlMessage::from_slice(&chunk.bytes) {
Ok(ControlMessage::Exit { code }) => {
let _ = exit_tx.send(Some(ExitOutcome::Exited(code)));
exit_resolved = true;
debug!("tty: exit chunk received, code={code}");
break;
}
Ok(other) => {
debug!("tty: ignoring non-exit control on STREAM_CTRL_OUT: {other:?}");
}
Err(e) => {
let _ = exit_tx.send(Some(ExitOutcome::MalformedExit(e.to_string())));
exit_resolved = true;
break;
}
},
STREAM_CTRL_IN | crate::wire::STREAM_STDIN => {
debug!(
"tty: ignoring client→server stream_type {} from server",
chunk.stream_type
);
}
other => {
debug!("tty: ignoring unknown stream_type {other}");
}
},
Err(RawError::ConnectionClosed) => {
debug!("tty: read pump: stream closed");
break;
}
Err(e) => {
warn!("tty: read pump: chunk read error: {e}");
break;
}
}
}
drop(stdout_tx);
drop(stderr_tx);
if !exit_resolved {
let _ = exit_tx.send(Some(ExitOutcome::NoExitChunk));
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::{MockBackend, TtyBackend};
use crate::negotiation::NegotiateRequest;
use alkcall::core::auth::Identity;
use alkcall::core::types::Connection;
use futures::stream::StreamExt;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::io::duplex;
fn test_negotiate(backend: &str) -> NegotiateRequest {
NegotiateRequest {
carriage: "raw".to_string(),
backend: backend.to_string(),
tty: None,
cmd: vec!["true".to_string()],
cwd: None,
env: HashMap::new(),
backend_params: serde_json::Map::new(),
}
}
async fn wire_session_and_server(
backend: Arc<dyn TtyBackend>,
identity: Option<Identity>,
) -> (TtySession, tokio::task::JoinHandle<()>) {
let mut backends: HashMap<String, Arc<dyn TtyBackend>> = HashMap::new();
backends.insert("mock".to_string(), backend);
let backends = Arc::new(backends);
let (client, server) = duplex(64 * 1024);
let (server_read, server_write) = tokio::io::split(server);
let server_task = tokio::spawn(async move {
crate::adapter::drive_session(server_write, server_read, backends, None, identity)
.await;
});
let client_conn = Connection::from_bidi(client, b"alk/tty".to_vec(), None);
let session = TtySession::connect_direct(client_conn, test_negotiate("mock"))
.await
.expect("connect_direct");
(session, server_task)
}
#[tokio::test]
async fn connect_direct_writes_negotiation_frame() {
let backend = Arc::new(MockBackend::with_exit_code(0));
let identity = Some(Identity {
id: "alice".to_string(),
scopes: vec![crate::adapter::TTY_OPEN_SCOPE.to_string()],
resources: HashMap::new(),
});
let (session, _server) = wire_session_and_server(backend, identity).await;
let code = tokio::time::timeout(std::time::Duration::from_secs(5), session.wait())
.await
.expect("wait didn't time out")
.expect("wait returns exit code");
assert_eq!(code, 0);
}
struct GatedBackend {
release: Arc<tokio::sync::Mutex<Option<tokio::sync::oneshot::Receiver<()>>>>,
}
#[async_trait::async_trait]
impl TtyBackend for GatedBackend {
async fn allocate(
&self,
_params: &crate::backend::TtyParams,
) -> Result<crate::backend::TtyHandle, crate::backend::TtyError> {
use crate::backend::{TtyControlHandle, TtyHandle};
let (_stdout_tx, stdout_rx) = mpsc::channel::<Bytes>(8);
let release = self.release.lock().await.take();
let stdout: Pin<Box<dyn Stream<Item = Bytes> + Send>> =
Box::pin(tokio_stream::wrappers::ReceiverStream::new(stdout_rx));
let stdin: Box<dyn AsyncWrite + Send + Unpin> = Box::new(tokio::io::sink());
let control = Some(TtyControlHandle::new(Arc::new(
crate::backend::MockControl::default(),
)));
let exit_code: crate::backend::BoxFuture<Result<i32, crate::backend::TtyError>> =
Box::pin(async move {
if let Some(rx) = release {
let _ = rx.await;
}
Ok(0)
});
Ok(TtyHandle {
stdin,
stdout,
stderr: None,
exit_code,
control,
})
}
}
async fn wire_gated_session() -> (
TtySession,
tokio::sync::oneshot::Sender<()>,
tokio::task::JoinHandle<()>,
) {
let (release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let backend = Arc::new(GatedBackend {
release: Arc::new(tokio::sync::Mutex::new(Some(release_rx))),
});
let identity = Some(Identity {
id: "alice".to_string(),
scopes: vec![crate::adapter::TTY_OPEN_SCOPE.to_string()],
resources: HashMap::new(),
});
let (session, server) = wire_session_and_server(backend, identity).await;
(session, release_tx, server)
}
#[tokio::test]
async fn send_stdin_round_trips_to_backend() {
let (session, release, _server) = wire_gated_session().await;
session
.send_stdin(Bytes::from_static(b"hello"))
.await
.expect("send_stdin");
session.close_stdin().await.expect("close_stdin");
release.send(()).unwrap();
let code = tokio::time::timeout(std::time::Duration::from_secs(5), session.wait())
.await
.expect("wait didn't time out")
.expect("wait returns exit code");
assert_eq!(code, 0);
}
#[tokio::test]
async fn resize_and_signal_dont_error() {
let (session, release, _server) = wire_gated_session().await;
session.resize(80, 24, 0, 0).await.expect("resize");
session.signal("INT").await.expect("signal");
release.send(()).unwrap();
let code = tokio::time::timeout(std::time::Duration::from_secs(5), session.wait())
.await
.expect("wait didn't time out")
.expect("wait returns exit code");
assert_eq!(code, 0);
}
#[tokio::test]
async fn recv_stdout_yields_backend_stdout() {
let backend = Arc::new(MockBackend::with_exit_code(0));
let identity = Some(Identity {
id: "alice".to_string(),
scopes: vec![crate::adapter::TTY_OPEN_SCOPE.to_string()],
resources: HashMap::new(),
});
let (session, _server) = wire_session_and_server(backend, identity).await;
let stdout = session.recv_stdout().await;
let collected: Vec<Bytes> = stdout.collect().await;
assert!(
collected.is_empty() || collected.iter().all(|b| b.is_empty()),
"mock backend produces no stdout, got {collected:?}"
);
}
#[tokio::test]
async fn wait_returns_no_exit_chunk_when_server_drops_without_exit() {
let (client, mut server) = duplex(64);
let server_handle = tokio::spawn(async move {
use tokio::io::AsyncReadExt;
let mut len_buf = [0u8; 4];
let _ = server.read_exact(&mut len_buf).await;
let len = u32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
let _ = server.read_exact(&mut body).await;
});
let client_conn = Connection::from_bidi(client, b"alk/tty".to_vec(), None);
let session = TtySession::connect_direct(client_conn, test_negotiate("mock"))
.await
.expect("connect_direct");
let result = tokio::time::timeout(std::time::Duration::from_secs(5), session.wait())
.await
.expect("wait didn't time out");
assert!(matches!(result, Err(TtySessionError::NoExitChunk)));
let _ = server_handle.await;
}
#[tokio::test]
async fn connect_direct_errors_when_stream_is_broken() {
let (client, server) = duplex(64);
drop(server);
let client_conn = Connection::from_bidi(client, b"alk/tty".to_vec(), None);
let result = TtySession::connect_direct(client_conn, test_negotiate("mock")).await;
assert!(
result.is_err(),
"construction should fail when the negotiation write hits a broken pipe"
);
}
struct EmittingBackend {
stdout: Vec<Bytes>,
stderr: Vec<Bytes>,
exit_code: i32,
}
#[async_trait::async_trait]
impl TtyBackend for EmittingBackend {
async fn allocate(
&self,
_params: &crate::backend::TtyParams,
) -> Result<crate::backend::TtyHandle, crate::backend::TtyError> {
use crate::backend::{TtyControlHandle, TtyHandle};
use tokio_stream::wrappers::ReceiverStream;
let (stdout_tx, stdout_rx) = mpsc::channel::<Bytes>(8);
let (stderr_tx, stderr_rx) = mpsc::channel::<Bytes>(8);
let (_stdin_tx, _stdin_rx) = mpsc::channel::<Bytes>(8);
let (exit_tx, exit_rx) =
tokio::sync::oneshot::channel::<Result<i32, crate::backend::TtyError>>();
let stdout = self.stdout.clone();
let stderr = self.stderr.clone();
let code = self.exit_code;
tokio::spawn(async move {
for b in stdout {
let _ = stdout_tx.send(b).await;
}
for b in stderr {
let _ = stderr_tx.send(b).await;
}
let _ = exit_tx.send(Ok(code));
});
let stdout: Pin<Box<dyn Stream<Item = Bytes> + Send>> =
Box::pin(ReceiverStream::new(stdout_rx));
let stderr: Option<Pin<Box<dyn Stream<Item = Bytes> + Send>>> =
Some(Box::pin(ReceiverStream::new(stderr_rx)));
let stdin: Box<dyn AsyncWrite + Send + Unpin> = Box::new(tokio::io::sink());
let control = Some(TtyControlHandle::new(Arc::new(
crate::backend::MockControl::default(),
)));
let exit_code: crate::backend::BoxFuture<Result<i32, crate::backend::TtyError>> =
Box::pin(async move {
exit_rx
.await
.map_err(|_| crate::backend::TtyError::WaitFailed {
message: "exit sender dropped".to_string(),
})
.and_then(|r| r)
});
Ok(TtyHandle {
stdin,
stdout,
stderr,
exit_code,
control,
})
}
}
#[tokio::test]
async fn recv_stdout_and_stderr_route_backend_data() {
let backend = Arc::new(EmittingBackend {
stdout: vec![Bytes::from_static(b"out1"), Bytes::from_static(b"out2")],
stderr: vec![Bytes::from_static(b"err1")],
exit_code: 0,
});
let identity = Some(Identity {
id: "alice".to_string(),
scopes: vec![crate::adapter::TTY_OPEN_SCOPE.to_string()],
resources: HashMap::new(),
});
let (session, _server) = wire_session_and_server(backend, identity).await;
let stdout = session.recv_stdout().await;
let collected: Vec<Bytes> = stdout.collect().await;
assert_eq!(
collected,
vec![Bytes::from_static(b"out1"), Bytes::from_static(b"out2")],
"stdout chunks should route to the stdout stream"
);
let stderr = session.recv_stderr().await.expect("stderr present");
let collected: Vec<Bytes> = stderr.collect().await;
assert_eq!(
collected,
vec![Bytes::from_static(b"err1")],
"stderr chunks should route to the stderr stream"
);
let code = tokio::time::timeout(std::time::Duration::from_secs(5), session.wait())
.await
.expect("wait didn't time out")
.expect("wait returns exit code");
assert_eq!(code, 0);
}
#[tokio::test]
async fn wait_returns_malformed_exit_chunk() {
let (client, mut server) = duplex(64 * 1024);
let server_handle = tokio::spawn(async move {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut len_buf = [0u8; 4];
let _ = server.read_exact(&mut len_buf).await;
let len = u32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
let _ = server.read_exact(&mut body).await;
let payload = br#"{"type":"exit","code":"not-a-number"}"#;
let mut header = [0u8; 5];
header[0] = crate::wire::STREAM_CTRL_OUT;
header[1..].copy_from_slice(&(payload.len() as u32).to_be_bytes());
let _ = server.write_all(&header).await;
let _ = server.write_all(payload).await;
let _ = server.flush().await;
});
let client_conn = Connection::from_bidi(client, b"alk/tty".to_vec(), None);
let session = TtySession::connect_direct(client_conn, test_negotiate("mock"))
.await
.expect("connect_direct");
let result = tokio::time::timeout(std::time::Duration::from_secs(5), session.wait())
.await
.expect("wait didn't time out");
assert!(
matches!(result, Err(TtySessionError::MalformedExitChunk(_))),
"expected MalformedExitChunk, got {result:?}"
);
let _ = server_handle.await;
}
#[tokio::test]
async fn read_pump_ignores_client_to_server_stream_types_from_server() {
let (client, mut server) = duplex(64 * 1024);
let server_handle = tokio::spawn(async move {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut len_buf = [0u8; 4];
let _ = server.read_exact(&mut len_buf).await;
let len = u32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
let _ = server.read_exact(&mut body).await;
async fn write_chunk(
server: &mut tokio::io::DuplexStream,
stream_type: u8,
payload: &[u8],
) {
let mut header = [0u8; 5];
header[0] = stream_type;
header[1..].copy_from_slice(&(payload.len() as u32).to_be_bytes());
server.write_all(&header).await.unwrap();
if !payload.is_empty() {
server.write_all(payload).await.unwrap();
}
}
write_chunk(&mut server, crate::wire::STREAM_STDOUT, b"out1").await;
write_chunk(
&mut server,
crate::wire::STREAM_STDIN,
b"server-must-not-send-stdin",
)
.await;
write_chunk(
&mut server,
crate::wire::STREAM_CTRL_IN,
br#"{"type":"resize"}"#,
)
.await;
write_chunk(&mut server, crate::wire::STREAM_STDOUT, b"out2").await;
let exit = br#"{"type":"exit","code":3}"#;
write_chunk(&mut server, crate::wire::STREAM_CTRL_OUT, exit).await;
let _ = server.flush().await;
});
let client_conn = Connection::from_bidi(client, b"alk/tty".to_vec(), None);
let session = TtySession::connect_direct(client_conn, test_negotiate("mock"))
.await
.expect("connect_direct");
let stdout = session.recv_stdout().await;
let collected: Vec<Bytes> = stdout.collect().await;
assert_eq!(
collected,
vec![Bytes::from_static(b"out1"), Bytes::from_static(b"out2")],
"stdout routing must be unaffected by the ignored chunks"
);
let code = tokio::time::timeout(std::time::Duration::from_secs(5), session.wait())
.await
.expect("wait didn't time out")
.expect("wait returns exit code");
assert_eq!(code, 3);
let _ = server_handle.await;
}
#[tokio::test]
async fn connect_direct_returns_negotiation_rejected() {
let (client, mut server) = duplex(64 * 1024);
let server_handle = tokio::spawn(async move {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut len_buf = [0u8; 4];
let _ = server.read_exact(&mut len_buf).await;
let len = u32::from_be_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
let _ = server.read_exact(&mut body).await;
let err_body = br#"{"error":"unknown_backend","backend":"nope"}"#;
let _ = server
.write_all(&(err_body.len() as u32).to_be_bytes())
.await;
let _ = server.write_all(err_body).await;
let _ = server.flush().await;
});
let client_conn = Connection::from_bidi(client, b"alk/tty".to_vec(), None);
let result = TtySession::connect_direct(client_conn, test_negotiate("mock")).await;
match result {
Err(TtySessionError::NegotiationRejected { error, fields }) => {
assert_eq!(error, "unknown_backend");
assert_eq!(fields.get("backend").map(String::as_str), Some("nope"));
}
Ok(_) => panic!("expected NegotiationRejected, got Ok(session)"),
Err(other) => panic!("expected NegotiationRejected, got {other:?}"),
}
let _ = server_handle.await;
}
use crate::testing::{tty_identity, wire_client_and_server};
fn test_open_params(backend: &str) -> serde_json::Value {
serde_json::json!({
"carriage": "raw",
"backend": backend,
"cmd": ["true"],
})
}
fn mock_backends(code: i32) -> Arc<HashMap<String, Arc<dyn TtyBackend>>> {
let mut backends: HashMap<String, Arc<dyn TtyBackend>> = HashMap::new();
backends.insert(
"mock".to_string(),
Arc::new(MockBackend::with_exit_code(code)),
);
Arc::new(backends)
}
#[tokio::test]
async fn open_via_channels_end_to_end_negotiates_and_waits() {
let client =
wire_client_and_server(mock_backends(0), None, Some(tty_identity("alice"))).await;
let session = tokio::time::timeout(
std::time::Duration::from_secs(10),
TtySession::open_via_channels(&client, test_open_params("mock")),
)
.await
.expect("open_via_channels timed out")
.expect("session opens");
let code = tokio::time::timeout(std::time::Duration::from_secs(10), session.wait())
.await
.expect("wait didn't time out")
.expect("wait returns exit code");
assert_eq!(code, 0);
}
#[tokio::test]
async fn open_via_channels_surfaces_negotiation_rejected() {
let client =
wire_client_and_server(mock_backends(0), None, Some(tty_identity("alice"))).await;
let result = tokio::time::timeout(
std::time::Duration::from_secs(10),
TtySession::open_via_channels(&client, test_open_params("nope")),
)
.await
.expect("open_via_channels timed out");
match result {
Err(TtySessionError::ChannelsOpen(
alkcall::channels::client::ChannelOpenError::CallFailed { error },
)) => {
assert_eq!(error.code, "channel:open_failed");
let details = error.details.expect("details carry the reason");
assert_eq!(details["reason"], "unknown_resource");
assert_eq!(details["message"], "unknown backend: nope");
}
Ok(_) => panic!("expected channel:open_failed, got Ok(session)"),
Err(other) => panic!("expected ChannelsOpen(channel:open_failed), got {other:?}"),
}
assert!(
client.manager().channel_ids().iter().all(|&id| id == 0),
"no data channel survives a failed establishment (channel 0 is the call channel)"
);
}
#[tokio::test]
async fn open_via_channels_fails_fast_on_schema_invalid_params() {
let client =
wire_client_and_server(mock_backends(0), None, Some(tty_identity("alice"))).await;
let result = tokio::time::timeout(
std::time::Duration::from_secs(10),
TtySession::open_via_channels(
&client,
serde_json::json!({ "carriage": "raw", "cmd": ["true"] }),
),
)
.await
.expect("open_via_channels timed out");
assert!(
matches!(result, Err(TtySessionError::InvalidParams(_))),
"schema-invalid params fail at the local NegotiateRequest parse (fail-fast, pre-open)"
);
}
#[tokio::test]
async fn open_via_channels_fails_fast_on_unparseable_params() {
let client =
wire_client_and_server(mock_backends(0), None, Some(tty_identity("alice"))).await;
let result = tokio::time::timeout(
std::time::Duration::from_secs(10),
TtySession::open_via_channels(
&client,
serde_json::json!({ "carriage": 42, "backend": "mock", "cmd": ["true"] }),
),
)
.await
.expect("open_via_channels timed out");
assert!(
matches!(result, Err(TtySessionError::InvalidParams(_))),
"unparseable params must fail before the open op"
);
}
#[tokio::test]
async fn open_via_channels_routes_backend_stdout_and_stderr() {
let mut backends: HashMap<String, Arc<dyn TtyBackend>> = HashMap::new();
backends.insert(
"mock".to_string(),
Arc::new(EmittingBackend {
stdout: vec![Bytes::from_static(b"ch-out")],
stderr: vec![Bytes::from_static(b"ch-err")],
exit_code: 3,
}),
);
let client =
wire_client_and_server(Arc::new(backends), None, Some(tty_identity("alice"))).await;
let session = tokio::time::timeout(
std::time::Duration::from_secs(10),
TtySession::open_via_channels(&client, test_open_params("mock")),
)
.await
.expect("open_via_channels timed out")
.expect("session opens");
let stdout = session.recv_stdout().await;
let collected: Vec<Bytes> = stdout.collect().await;
assert_eq!(
collected,
vec![Bytes::from_static(b"ch-out")],
"stdout should route through the channels data plane"
);
let stderr = session.recv_stderr().await.expect("stderr present");
let collected: Vec<Bytes> = stderr.collect().await;
assert_eq!(
collected,
vec![Bytes::from_static(b"ch-err")],
"stderr should route through the channels data plane"
);
let code = tokio::time::timeout(std::time::Duration::from_secs(10), session.wait())
.await
.expect("wait didn't time out")
.expect("wait returns exit code");
assert_eq!(code, 3);
}
}