use rho_sdk::CancellationToken;
use tokio::io::{AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::process::ChildStdin;
use tokio::sync::{mpsc, oneshot};
use super::{
child::OwnedChild,
line_decoder::{claude_ndjson_line_decoder, LineDecodeError},
messaging,
stream::{StreamEffect, StreamMapper, TerminalResult},
};
const MAX_STDERR_BYTES: usize = 8 * 1024;
const READ_CHUNK_BYTES: usize = 8 * 1024;
const STDIN_BROKEN_PIPE: &str = "claude code: stdin closed by child (broken pipe)";
pub(crate) enum DrainInput {
Text { prompt: String },
StreamJson {
initial_prompt: String,
parent_messages: Option<messaging::ClaudeMessageInbox>,
},
}
pub(crate) enum DrainEnd {
Cancelled,
StdinFailed(String),
StreamFailed(String),
Exited(std::io::Result<std::process::ExitStatus>),
}
pub(crate) struct Drained {
pub(crate) terminal: Option<TerminalResult>,
pub(crate) stderr: String,
pub(crate) end: DrainEnd,
}
pub(crate) async fn drain_child(
child: &mut OwnedChild,
input: DrainInput,
cancellation: &CancellationToken,
on_effect: &mut (dyn FnMut(StreamEffect) + Send),
) -> Drained {
let Some(stdout) = child.stdout() else {
return Drained {
terminal: None,
stderr: String::new(),
end: DrainEnd::StreamFailed("claude code: child stdout was not captured".into()),
};
};
let Some(stdin) = child.stdin() else {
return Drained {
terminal: None,
stderr: String::new(),
end: DrainEnd::StdinFailed("claude code: child stdin was not captured".into()),
};
};
let (close_tx, stdin_write) = match input {
DrainInput::Text { prompt } => {
let write = tokio::spawn(async move { write_text_stdin(stdin, prompt).await });
(None, write)
}
DrainInput::StreamJson {
initial_prompt,
parent_messages,
} => {
let (close_tx, close_rx) = oneshot::channel::<()>();
let write = tokio::spawn(async move {
write_stream_json_stdin(stdin, initial_prompt, parent_messages, close_rx).await
});
(Some(close_tx), write)
}
};
tokio::pin!(stdin_write);
let stderr = child.stderr();
let read_stderr = async move {
let mut tail = StderrTail::default();
let Some(mut stderr) = stderr else {
return tail;
};
let mut chunk = vec![0_u8; READ_CHUNK_BYTES];
loop {
match stderr.read(&mut chunk).await {
Ok(0) | Err(_) => return tail,
Ok(count) => tail.push(&chunk[..count]),
}
}
};
tokio::pin!(read_stderr);
let mut stdout = BufReader::new(stdout);
let mut decoder = claude_ndjson_line_decoder();
let mut mapper = StreamMapper::new();
let mut terminal: Option<TerminalResult> = None;
let mut stderr_text = String::new();
let mut stderr_done = false;
let mut stdout_done = false;
let mut stdin_done = false;
let mut exit_result = None;
let mut close_tx = close_tx;
let mut chunk = vec![0_u8; READ_CHUNK_BYTES];
let early_end: Option<DrainEnd> = loop {
if stdin_done && stdout_done && stderr_done && exit_result.is_some() {
break None;
}
tokio::select! {
biased;
() = cancellation.cancelled() => {
break Some(DrainEnd::Cancelled);
}
result = &mut stdin_write, if !stdin_done => {
stdin_done = true;
if let Ok(Err(error)) = result {
if error != STDIN_BROKEN_PIPE {
break Some(DrainEnd::StdinFailed(error));
}
}
}
captured = &mut read_stderr, if !stderr_done => {
stderr_done = true;
stderr_text = captured.finish();
}
status = child.wait(), if exit_result.is_none() => {
exit_result = Some(status);
drop(close_tx.take());
}
read = stdout.read(&mut chunk), if !stdout_done => {
match read {
Ok(0) => stdout_done = true,
Ok(count) => {
decoder.push(&chunk[..count]);
let mut decode_error = None;
loop {
match decoder.next_line() {
Ok(Some(line)) => {
let line = line.to_string();
for effect in mapper.push_line(&line) {
if let StreamEffect::Terminal(result) = &effect {
terminal = Some(result.clone());
if let Some(tx) = close_tx.take() {
let _ = tx.send(());
}
}
on_effect(effect);
}
}
Ok(None) => break,
Err(error) => {
decode_error = Some(format_line_error(&error));
break;
}
}
}
if let Some(error) = decode_error {
break Some(DrainEnd::StreamFailed(error));
}
}
Err(error) => {
break Some(DrainEnd::StreamFailed(format!(
"claude code: failed reading stdout: {error}"
)));
}
}
}
}
};
let end = match early_end {
Some(end) => end,
None => match decoder.finish() {
Err(error) => DrainEnd::StreamFailed(format_line_error(&error)),
Ok(tail) => {
if let Some(line) = tail {
for effect in mapper.push_line(line) {
if let StreamEffect::Terminal(result) = &effect {
terminal = Some(result.clone());
}
on_effect(effect);
}
}
DrainEnd::Exited(exit_result.expect("completed drain reaped the child"))
}
},
};
Drained {
terminal,
stderr: stderr_text,
end,
}
}
async fn write_text_stdin(mut stdin: ChildStdin, prompt: String) -> Result<(), String> {
write_all(&mut stdin, prompt.as_bytes()).await?;
shutdown_stdin(&mut stdin).await
}
async fn write_stream_json_stdin(
mut stdin: ChildStdin,
initial_prompt: String,
mut parent_messages: Option<messaging::ClaudeMessageInbox>,
close_rx: oneshot::Receiver<()>,
) -> Result<(), String> {
write_all(
&mut stdin,
messaging::encode_user_turn(&initial_prompt).as_bytes(),
)
.await?;
let mut close_rx = Some(close_rx);
loop {
if let Some(inbox) = parent_messages.as_mut() {
match inbox.try_recv() {
Ok(text) => {
write_parent_turn(&mut stdin, &text).await?;
continue;
}
Err(mpsc::error::TryRecvError::Empty) => {}
Err(mpsc::error::TryRecvError::Disconnected) => {
parent_messages = None;
}
}
}
let Some(close) = close_rx.as_mut() else {
break;
};
tokio::select! {
biased;
result = close => {
let _ = result;
close_rx = None;
if let Some(inbox) = parent_messages.as_ref() {
inbox.seal();
}
}
maybe_text = recv_parent(&mut parent_messages), if parent_messages.is_some() => {
if let Some(text) = maybe_text {
write_parent_turn(&mut stdin, &text).await?;
}
}
}
}
if let Some(mut inbox) = parent_messages.take() {
inbox.seal();
while let Some(text) = inbox.recv().await {
write_parent_turn(&mut stdin, &text).await?;
}
}
shutdown_stdin(&mut stdin).await
}
async fn write_parent_turn(stdin: &mut ChildStdin, text: &str) -> Result<(), String> {
write_all(
stdin,
messaging::encode_user_turn(&messaging::frame_parent_message(text)).as_bytes(),
)
.await
}
async fn write_all(stdin: &mut ChildStdin, mut bytes: &[u8]) -> Result<(), String> {
while !bytes.is_empty() {
match stdin.write(bytes).await {
Ok(0) => {
return Err("claude code: failed to write prompt to stdin: wrote 0 bytes".into());
}
Ok(count) => bytes = &bytes[count..],
Err(error) => return Err(map_stdin_io_error(error)),
}
}
Ok(())
}
async fn shutdown_stdin(stdin: &mut ChildStdin) -> Result<(), String> {
match stdin.shutdown().await {
Ok(()) => Ok(()),
Err(error) => Err(map_stdin_io_error(error)),
}
}
fn map_stdin_io_error(error: std::io::Error) -> String {
if error.kind() == std::io::ErrorKind::BrokenPipe {
STDIN_BROKEN_PIPE.into()
} else {
format!("claude code: failed to write prompt to stdin: {error}")
}
}
async fn recv_parent(
parent_messages: &mut Option<messaging::ClaudeMessageInbox>,
) -> Option<String> {
let Some(inbox) = parent_messages.as_mut() else {
std::future::pending::<()>().await;
unreachable!("pending future resolved");
};
match inbox.recv().await {
Some(text) => Some(text),
None => {
*parent_messages = None;
None
}
}
}
pub(crate) fn format_line_error(error: &LineDecodeError) -> String {
match error {
LineDecodeError::InvalidUtf8(_) => {
format!("claude code: malformed UTF-8 on stream-json stdout: {error}")
}
LineDecodeError::LineTooLong { .. } => {
format!("claude code: oversize stream-json line: {error}")
}
}
}
#[derive(Default)]
struct StderrTail {
bytes: Vec<u8>,
elided: bool,
}
impl StderrTail {
fn push(&mut self, chunk: &[u8]) {
self.bytes.extend_from_slice(chunk);
if self.bytes.len() <= MAX_STDERR_BYTES {
return;
}
let cut = ceil_utf8_boundary(&self.bytes, self.bytes.len() - MAX_STDERR_BYTES);
self.bytes.drain(..cut);
self.elided = true;
}
fn finish(self) -> String {
let text = String::from_utf8_lossy(&self.bytes);
let trimmed = text.trim();
if self.elided {
format!("{}{trimmed}", rho_sdk::ELLIPSIS)
} else {
trimmed.to_string()
}
}
}
fn ceil_utf8_boundary(bytes: &[u8], index: usize) -> usize {
let mut index = index.min(bytes.len());
while index < bytes.len() && bytes[index] & 0b1100_0000 == 0b1000_0000 {
index += 1;
}
index
}
#[cfg(test)]
#[path = "drain_tests.rs"]
mod tests;