use std::collections::{HashMap, HashSet};
use std::path::PathBuf;
use std::process::Stdio;
use std::sync::{Arc, Mutex as StdMutex};
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt};
use tokio::process::{Child, ChildStdin, Command};
use tokio::sync::{Mutex, OwnedSemaphorePermit, Semaphore, mpsc, oneshot, watch};
use tokio::time::{Instant, timeout, timeout_at};
use crate::domain::errors::{AgentError, AgentResult, ErrorCode};
use crate::domain::pi_rpc::{PiRpcCommand, PiRpcEvent, PiRpcOutput, PiRpcResponse};
use crate::infrastructure::pi_rpc_framing::PiRpcDecoder;
const MAX_RECORD_BYTES: usize = 2 * 1024 * 1024;
const READ_CHUNK_BYTES: usize = 32 * 1024;
const STDERR_CAPACITY_BYTES: usize = 64 * 1024;
const EVENT_QUEUE_CAPACITY: usize = 1024;
pub(crate) const EVENT_QUEUE_MAX_BYTES: usize = 16 * 1024 * 1024;
const WRITER_QUEUE_CAPACITY: usize = 16;
const MAX_COMMAND_RECORD_BYTES: usize = MAX_RECORD_BYTES;
const SHUTDOWN_GRACE: Duration = Duration::from_millis(250);
pub(crate) struct PiProcessOptions {
pub(crate) command: Vec<String>,
pub(crate) cwd: PathBuf,
pub(crate) session: PiSessionMode,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum PiSessionMode {
New,
Resume(PathBuf),
Ephemeral,
}
pub(crate) struct PiRpcProcess {
state: Arc<ProcessState>,
writer_tx: mpsc::Sender<WriteJob>,
supervisor_tx: mpsc::UnboundedSender<SupervisorCommand>,
pid: Option<u32>,
}
pub(crate) struct SpawnedPiRpcProcess {
pub(crate) process: PiRpcProcess,
pub(crate) events: PiRpcEventReceiver,
}
pub(crate) struct PiRpcEventReceiver {
receiver: mpsc::Receiver<EventEnvelope>,
#[cfg(test)]
budget: Arc<Semaphore>,
}
pub(crate) struct PiRpcSequencedEvent {
pub(crate) sequence: u64,
pub(crate) event: PiRpcEvent,
}
pub(crate) struct PiRpcReply {
pub(crate) response: PiRpcResponse,
pub(crate) event_barrier: u64,
}
struct ProcessState {
core: StdMutex<ProcessCore>,
failure_tx: watch::Sender<Option<FailureReason>>,
stderr: Mutex<Vec<u8>>,
}
struct ProcessCore {
pending: HashMap<String, PendingRequest>,
used_ids: HashSet<String>,
failure: Option<FailureReason>,
}
struct PendingRequest {
command: &'static str,
sender: oneshot::Sender<AgentResult<PiRpcReply>>,
}
struct PendingGuard {
state: Arc<ProcessState>,
id: String,
}
struct WriteJob {
record: Vec<u8>,
acknowledgement: oneshot::Sender<AgentResult<()>>,
}
struct EventEnvelope {
sequence: u64,
event: PiRpcEvent,
_byte_permit: OwnedSemaphorePermit,
}
enum SupervisorCommand {
Shutdown(oneshot::Sender<AgentResult<()>>),
Force,
}
enum SupervisorTrigger {
ChildExited(std::io::Result<std::process::ExitStatus>),
StdoutFinished(Result<Option<FailureReason>, tokio::task::JoinError>),
WriterStopped,
Command(Option<SupervisorCommand>),
}
#[derive(Clone)]
struct FailureReason {
code: ErrorCode,
message: String,
}
impl FailureReason {
fn error(&self) -> AgentError {
AgentError::new(self.code, self.message.clone())
}
}
impl PiRpcEventReceiver {
pub(crate) async fn recv(&mut self) -> Option<PiRpcEvent> {
self.recv_sequenced().await.map(|event| event.event)
}
pub(crate) async fn recv_sequenced(&mut self) -> Option<PiRpcSequencedEvent> {
let envelope = self.receiver.recv().await?;
let EventEnvelope {
sequence,
event,
_byte_permit,
} = envelope;
drop(_byte_permit);
Some(PiRpcSequencedEvent { sequence, event })
}
#[cfg(test)]
pub(crate) fn queued_bytes(&self) -> usize {
EVENT_QUEUE_MAX_BYTES - self.budget.available_permits()
}
}
impl PiRpcProcess {
pub(crate) async fn spawn(options: PiProcessOptions) -> AgentResult<SpawnedPiRpcProcess> {
let Some((program, prefix_args)) = options.command.split_first() else {
return Err(process_error("Pi RPC command cannot be empty"));
};
let mut command = Command::new(program);
command
.args(prefix_args)
.args(["--mode", "rpc"])
.current_dir(&options.cwd)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.kill_on_drop(true);
match &options.session {
PiSessionMode::New => {}
PiSessionMode::Resume(session) => {
command.arg("--session").arg(session);
}
PiSessionMode::Ephemeral => {
command.arg("--no-session");
}
}
let mut child = command
.spawn()
.map_err(|error| process_error(format!("failed to start Pi RPC process: {error}")))?;
let pid = child.id();
let stdin = child
.stdin
.take()
.ok_or_else(|| process_error("Pi RPC process stdin was not piped"))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| process_error("Pi RPC process stdout was not piped"))?;
let stderr = child
.stderr
.take()
.ok_or_else(|| process_error("Pi RPC process stderr was not piped"))?;
let (event_tx, event_rx) = mpsc::channel(EVENT_QUEUE_CAPACITY);
let event_budget = Arc::new(Semaphore::new(EVENT_QUEUE_MAX_BYTES));
let (failure_tx, _) = watch::channel(None);
let state = Arc::new(ProcessState {
core: StdMutex::new(ProcessCore {
pending: HashMap::new(),
used_ids: HashSet::new(),
failure: None,
}),
failure_tx,
stderr: Mutex::new(Vec::new()),
});
let (fatal_tx, fatal_rx) = mpsc::unbounded_channel();
let (writer_tx, writer_rx) = mpsc::channel(WRITER_QUEUE_CAPACITY);
let writer_task = tokio::spawn(run_writer(stdin, writer_rx, state.clone(), fatal_tx));
let stdout_task = tokio::spawn(read_stdout(
stdout,
state.clone(),
event_tx,
event_budget.clone(),
));
let stderr_task = tokio::spawn(read_stderr(stderr, state.clone()));
let (supervisor_tx, supervisor_rx) = mpsc::unbounded_channel();
let supervisor_state = state.clone();
tokio::spawn(supervise_child(
child,
supervisor_state,
fatal_rx,
supervisor_rx,
writer_task,
stdout_task,
stderr_task,
));
Ok(SpawnedPiRpcProcess {
process: Self {
state,
writer_tx,
supervisor_tx,
pid,
},
events: PiRpcEventReceiver {
receiver: event_rx,
#[cfg(test)]
budget: event_budget,
},
})
}
pub(crate) async fn request(
&self,
command: PiRpcCommand,
request_timeout: Duration,
) -> AgentResult<PiRpcResponse> {
self.request_with_barrier(command, request_timeout)
.await
.map(|reply| reply.response)
}
pub(crate) async fn request_with_barrier(
&self,
command: PiRpcCommand,
request_timeout: Duration,
) -> AgentResult<PiRpcReply> {
let deadline = Instant::now() + request_timeout;
if !command.expects_response() {
return Err(invalid_usage(
"Pi RPC request called with a fire-and-forget command",
));
}
let id = command.id().to_owned();
let expected_command = command_name(&command);
let (pending, receiver) = self.state.register_pending(id.clone(), expected_command)?;
let record = serialize_command(&command)?;
let timeout_message = format!("Pi RPC {expected_command} request {id} timed out");
let write_acknowledgement = self
.enqueue_record(record, deadline, &timeout_message)
.await?;
let timeout_error = || process_error(timeout_message.clone());
match timeout_at(deadline, write_acknowledgement).await {
Ok(Ok(result)) => result?,
Ok(Err(_)) => return Err(self.state.current_error().unwrap_or_else(timeout_error)),
Err(_) => return Err(timeout_error()),
}
let result = match timeout_at(deadline, receiver).await {
Ok(Ok(result)) => result,
Ok(Err(_)) => Err(self.state.current_error().unwrap_or_else(|| {
process_error("Pi RPC response waiter closed without a result")
})),
Err(_) => Err(timeout_error()),
};
drop(pending);
result
}
pub(crate) async fn send(
&self,
command: PiRpcCommand,
operation_timeout: Duration,
) -> AgentResult<()> {
let deadline = Instant::now() + operation_timeout;
if command.expects_response() {
return Err(invalid_usage(
"Pi RPC send called with a response-producing command",
));
}
self.state.reserve_id(command.id())?;
let command_name = command_name(&command);
let id = command.id().to_owned();
let timeout_message = format!("Pi RPC {command_name} send {id} timed out");
let record = serialize_command(&command)?;
let acknowledgement = self
.enqueue_record(record, deadline, &timeout_message)
.await?;
match timeout_at(deadline, acknowledgement).await {
Ok(Ok(result)) => result,
Err(_) => Err(process_error(timeout_message)),
Ok(Err(_)) => Err(self.state.current_error().unwrap_or_else(|| {
process_error("Pi RPC writer closed without acknowledging the command")
})),
}
}
#[cfg(test)]
pub(crate) async fn stderr_snapshot(&self) -> String {
String::from_utf8_lossy(&self.state.stderr.lock().await).into_owned()
}
pub(crate) async fn shutdown(&self) -> AgentResult<()> {
let (sender, receiver) = oneshot::channel();
if self
.supervisor_tx
.send(SupervisorCommand::Shutdown(sender))
.is_err()
{
return Ok(());
}
match receiver.await {
Ok(result) => result,
Err(_) => Ok(()),
}
}
pub(crate) fn pid(&self) -> Option<u32> {
self.pid
}
#[cfg(test)]
pub(crate) fn pending_request_count(&self) -> usize {
self.state.pending_request_count()
}
async fn enqueue_record(
&self,
record: Vec<u8>,
deadline: Instant,
timeout_message: &str,
) -> AgentResult<oneshot::Receiver<AgentResult<()>>> {
self.state.ensure_running()?;
let (acknowledgement, receiver) = oneshot::channel();
let job = WriteJob {
record,
acknowledgement,
};
match timeout_at(deadline, self.writer_tx.send(job)).await {
Ok(Ok(())) => Ok(receiver),
Ok(Err(_)) => Err(self
.state
.current_error()
.unwrap_or_else(|| process_error("Pi RPC writer is no longer accepting commands"))),
Err(_) => Err(process_error(timeout_message.to_owned())),
}
}
}
impl Drop for PiRpcProcess {
fn drop(&mut self) {
self.state.fail(FailureReason {
code: ErrorCode::BackendDisconnected,
message: "Pi RPC process was dropped".into(),
});
let _ = self.supervisor_tx.send(SupervisorCommand::Force);
}
}
impl ProcessState {
fn register_pending(
self: &Arc<Self>,
id: String,
command: &'static str,
) -> AgentResult<(PendingGuard, oneshot::Receiver<AgentResult<PiRpcReply>>)> {
let (sender, receiver) = oneshot::channel();
let mut core = self.core.lock().unwrap_or_else(|error| error.into_inner());
if let Some(failure) = &core.failure {
return Err(failure.error());
}
if !core.used_ids.insert(id.clone()) {
return Err(invalid_usage(format!(
"Pi RPC request id is already used: {id}"
)));
}
core.pending
.insert(id.clone(), PendingRequest { command, sender });
drop(core);
Ok((
PendingGuard {
state: self.clone(),
id,
},
receiver,
))
}
fn reserve_id(&self, id: &str) -> AgentResult<()> {
let mut core = self.core.lock().unwrap_or_else(|error| error.into_inner());
if let Some(failure) = &core.failure {
return Err(failure.error());
}
if !core.used_ids.insert(id.to_owned()) {
return Err(invalid_usage(format!(
"Pi RPC request id is already used: {id}"
)));
}
Ok(())
}
fn remove_pending(&self, id: &str) {
self.core
.lock()
.unwrap_or_else(|error| error.into_inner())
.pending
.remove(id);
}
fn ensure_running(&self) -> AgentResult<()> {
match &self
.core
.lock()
.unwrap_or_else(|error| error.into_inner())
.failure
{
Some(failure) => Err(failure.error()),
None => Ok(()),
}
}
fn current_error(&self) -> Option<AgentError> {
self.core
.lock()
.unwrap_or_else(|error| error.into_inner())
.failure
.as_ref()
.map(FailureReason::error)
}
#[cfg(test)]
fn pending_request_count(&self) -> usize {
self.core
.lock()
.unwrap_or_else(|error| error.into_inner())
.pending
.len()
}
fn complete_response(&self, response: PiRpcResponse, event_barrier: u64) -> AgentResult<()> {
let Some(id) = response.id().map(str::to_owned) else {
return Err(invalid_protocol(
"Pi RPC response is missing its request id",
));
};
let pending = self
.core
.lock()
.unwrap_or_else(|error| error.into_inner())
.pending
.remove(&id);
let Some(pending) = pending else {
return Ok(());
};
if pending.command != response.command() {
let error = invalid_protocol(format!(
"Pi RPC response {id} was for {}, expected {}",
response.command(),
pending.command
));
let _ = pending.sender.send(Err(invalid_protocol(error.message())));
return Err(error);
}
let _ = pending.sender.send(Ok(PiRpcReply {
response,
event_barrier,
}));
Ok(())
}
fn fail(&self, reason: FailureReason) {
let pending = {
let mut core = self.core.lock().unwrap_or_else(|error| error.into_inner());
if core.failure.is_some() {
None
} else {
core.failure = Some(reason.clone());
Some(std::mem::take(&mut core.pending))
}
};
if let Some(pending) = pending {
self.failure_tx.send_replace(Some(reason.clone()));
for request in pending.into_values() {
let _ = request.sender.send(Err(reason.error()));
}
}
}
}
impl Drop for PendingGuard {
fn drop(&mut self) {
self.state.remove_pending(&self.id);
}
}
async fn wait_for_failure(failure_rx: &mut watch::Receiver<Option<FailureReason>>) -> AgentError {
loop {
if let Some(reason) = failure_rx.borrow_and_update().clone() {
return reason.error();
}
if failure_rx.changed().await.is_err() {
return process_error("Pi RPC process failure channel closed");
}
}
}
async fn write_record(stdin: &mut ChildStdin, record: &[u8]) -> AgentResult<()> {
stdin
.write_all(record)
.await
.map_err(|error| process_error(format!("failed to write Pi RPC command: {error}")))?;
stdin
.flush()
.await
.map_err(|error| process_error(format!("failed to flush Pi RPC command: {error}")))
}
fn serialize_command(command: &PiRpcCommand) -> AgentResult<Vec<u8>> {
let mut record = serde_json::to_vec(command)
.map_err(|error| invalid_usage(format!("failed to serialize Pi RPC command: {error}")))?;
if record.len() >= MAX_COMMAND_RECORD_BYTES {
return Err(invalid_usage(
"serialized Pi RPC command exceeds the 2 MiB record limit including its LF terminator",
));
}
record.push(b'\n');
Ok(record)
}
async fn run_writer(
mut stdin: ChildStdin,
mut writer_rx: mpsc::Receiver<WriteJob>,
state: Arc<ProcessState>,
fatal_tx: mpsc::UnboundedSender<()>,
) {
let mut failure_rx = state.failure_tx.subscribe();
loop {
let mut job = tokio::select! {
biased;
_ = wait_for_failure(&mut failure_rx) => return,
job = writer_rx.recv() => match job {
Some(job) => job,
None => return,
},
};
if job.acknowledgement.is_closed() {
continue;
}
let result = tokio::select! {
biased;
error = wait_for_failure(&mut failure_rx) => {
let _ = job.acknowledgement.send(Err(error));
return;
}
result = write_record(&mut stdin, &job.record) => result,
_ = job.acknowledgement.closed() => {
let error = process_error(
"Pi RPC command caller cancelled while its record was being written",
);
state.fail(reason_from(&error));
let _ = fatal_tx.send(());
return;
}
};
match result {
Ok(()) => {
let _ = job.acknowledgement.send(Ok(()));
}
Err(error) => {
state.fail(reason_from(&error));
let _ = job.acknowledgement.send(Err(error));
let _ = fatal_tx.send(());
return;
}
}
}
}
async fn read_stdout<R>(
mut stdout: R,
state: Arc<ProcessState>,
event_tx: mpsc::Sender<EventEnvelope>,
event_budget: Arc<Semaphore>,
) -> Option<FailureReason>
where
R: AsyncRead + Unpin,
{
let mut decoder = PiRpcDecoder::new(MAX_RECORD_BYTES);
let mut chunk = vec![0; READ_CHUNK_BYTES];
let mut event_sequence = 0_u64;
loop {
match stdout.read(&mut chunk).await {
Ok(0) => match decoder.finish() {
Ok(_) => return None,
Err(error) => return Some(reason_from(&error)),
},
Ok(read) => match decoder.push_sized(&chunk[..read]) {
Ok(outputs) => {
let mut protocol_error = None;
for output in outputs {
match output.output {
PiRpcOutput::Response(response) => {
if let Err(error) =
state.complete_response(response, event_sequence)
{
protocol_error = Some(reason_from(&error));
break;
}
}
PiRpcOutput::Event(event) => {
event_sequence = event_sequence.saturating_add(1);
let permits = match u32::try_from(output.json_bytes) {
Ok(permits) => permits,
Err(_) => {
return Some(FailureReason {
code: ErrorCode::InvalidMessage,
message: "Pi RPC event is too large to budget".into(),
});
}
};
let byte_permit = match event_budget
.clone()
.try_acquire_many_owned(permits)
{
Ok(permit) => permit,
Err(tokio::sync::TryAcquireError::NoPermits) => {
return Some(FailureReason {
code: ErrorCode::InvalidMessage,
message: format!(
"Pi RPC event byte budget of {EVENT_QUEUE_MAX_BYTES} bytes was exceeded"
),
});
}
Err(tokio::sync::TryAcquireError::Closed) => {
return Some(FailureReason {
code: ErrorCode::BackendDisconnected,
message: "Pi RPC event byte budget is closed".into(),
});
}
};
if let Err(error) = enqueue_event_with_one_fair_retry(
&event_tx,
EventEnvelope {
sequence: event_sequence,
event,
_byte_permit: byte_permit,
},
)
.await
{
return Some(error);
}
}
}
}
if let Some(error) = protocol_error {
return Some(error);
}
}
Err(error) => return Some(reason_from(&error)),
},
Err(error) => {
return Some(FailureReason {
code: ErrorCode::BackendDisconnected,
message: format!("failed to read Pi RPC stdout: {error}"),
});
}
}
}
}
async fn enqueue_event_with_one_fair_retry(
event_tx: &mpsc::Sender<EventEnvelope>,
event: EventEnvelope,
) -> Result<(), FailureReason> {
let event = match event_tx.try_send(event) {
Ok(()) => return Ok(()),
Err(mpsc::error::TrySendError::Full(event)) => event,
Err(mpsc::error::TrySendError::Closed(_)) => return Err(event_consumer_closed()),
};
tokio::task::yield_now().await;
match event_tx.try_send(event) {
Ok(()) => Ok(()),
Err(mpsc::error::TrySendError::Full(_)) => Err(event_queue_full()),
Err(mpsc::error::TrySendError::Closed(_)) => Err(event_consumer_closed()),
}
}
fn event_queue_full() -> FailureReason {
FailureReason {
code: ErrorCode::InvalidMessage,
message: format!("Pi RPC event queue exceeded its {EVENT_QUEUE_CAPACITY}-event capacity"),
}
}
fn event_consumer_closed() -> FailureReason {
FailureReason {
code: ErrorCode::BackendDisconnected,
message: "Pi RPC event consumer is closed".into(),
}
}
async fn read_stderr<R>(mut stderr: R, state: Arc<ProcessState>)
where
R: AsyncRead + Unpin,
{
let mut chunk = vec![0; READ_CHUNK_BYTES];
loop {
match stderr.read(&mut chunk).await {
Ok(0) | Err(_) => return,
Ok(read) => append_bounded_stderr(&state.stderr, &chunk[..read]).await,
}
}
}
async fn append_bounded_stderr(stderr: &Mutex<Vec<u8>>, bytes: &[u8]) {
let mut buffer = stderr.lock().await;
if bytes.len() >= STDERR_CAPACITY_BYTES {
buffer.clear();
buffer.extend_from_slice(&bytes[bytes.len() - STDERR_CAPACITY_BYTES..]);
return;
}
let excess = buffer
.len()
.saturating_add(bytes.len())
.saturating_sub(STDERR_CAPACITY_BYTES);
if excess > 0 {
buffer.drain(..excess);
}
buffer.extend_from_slice(bytes);
}
async fn supervise_child(
mut child: Child,
state: Arc<ProcessState>,
mut fatal_rx: mpsc::UnboundedReceiver<()>,
mut supervisor_rx: mpsc::UnboundedReceiver<SupervisorCommand>,
mut writer_task: tokio::task::JoinHandle<()>,
mut stdout_task: tokio::task::JoinHandle<Option<FailureReason>>,
mut stderr_task: tokio::task::JoinHandle<()>,
) {
let mut shutdown_response = None;
let trigger = tokio::select! {
status = child.wait() => SupervisorTrigger::ChildExited(status),
outcome = &mut stdout_task => SupervisorTrigger::StdoutFinished(outcome),
_ = fatal_rx.recv() => SupervisorTrigger::WriterStopped,
command = supervisor_rx.recv() => SupervisorTrigger::Command(command),
};
match trigger {
SupervisorTrigger::ChildExited(status) => {
let stdout_failure = finish_stdout_task(&mut stdout_task).await;
state.fail(stdout_failure.unwrap_or_else(|| exit_failure(status)));
}
SupervisorTrigger::StdoutFinished(outcome) => {
if let Some(failure) = classify_stdout_result(outcome) {
state.fail(failure);
force_kill_and_wait(&mut child).await;
} else {
match timeout(SHUTDOWN_GRACE, child.wait()).await {
Ok(status) => state.fail(exit_failure(status)),
Err(_) => {
state.fail(FailureReason {
code: ErrorCode::BackendDisconnected,
message: "Pi RPC stdout closed".into(),
});
force_kill_and_wait(&mut child).await;
}
}
}
}
SupervisorTrigger::WriterStopped => {
state.fail(FailureReason {
code: ErrorCode::BackendDisconnected,
message: "Pi RPC writer stopped unexpectedly".into(),
});
force_kill_and_wait(&mut child).await;
let _ = finish_stdout_task(&mut stdout_task).await;
}
SupervisorTrigger::Command(command) => match command.unwrap_or(SupervisorCommand::Force) {
SupervisorCommand::Shutdown(response) => {
shutdown_response = Some(response);
state.fail(FailureReason {
code: ErrorCode::BackendDisconnected,
message: "Pi RPC process is shutting down".into(),
});
match timeout(SHUTDOWN_GRACE, child.wait()).await {
Ok(Ok(_)) => {}
Ok(Err(_)) | Err(_) => force_kill_and_wait(&mut child).await,
}
let _ = finish_stdout_task(&mut stdout_task).await;
}
SupervisorCommand::Force => {
state.fail(FailureReason {
code: ErrorCode::BackendDisconnected,
message: "Pi RPC process was dropped".into(),
});
force_kill_and_wait(&mut child).await;
let _ = finish_stdout_task(&mut stdout_task).await;
}
},
}
finish_unit_task(&mut writer_task).await;
finish_unit_task(&mut stderr_task).await;
if let Some(response) = shutdown_response {
let _ = response.send(Ok(()));
}
}
async fn finish_stdout_task(
task: &mut tokio::task::JoinHandle<Option<FailureReason>>,
) -> Option<FailureReason> {
match timeout(SHUTDOWN_GRACE, &mut *task).await {
Ok(result) => classify_stdout_result(result),
Err(_) => {
task.abort();
let _ = task.await;
None
}
}
}
async fn finish_unit_task(task: &mut tokio::task::JoinHandle<()>) {
if timeout(SHUTDOWN_GRACE, &mut *task).await.is_err() {
task.abort();
let _ = task.await;
}
}
fn classify_stdout_result(
result: Result<Option<FailureReason>, tokio::task::JoinError>,
) -> Option<FailureReason> {
match result {
Ok(failure) => failure,
Err(error) => Some(FailureReason {
code: ErrorCode::BackendDisconnected,
message: format!("Pi RPC stdout reader failed: {error}"),
}),
}
}
fn exit_failure(status: std::io::Result<std::process::ExitStatus>) -> FailureReason {
let message = match status {
Ok(status) => format!("Pi RPC process exited unexpectedly with {status}"),
Err(error) => format!("failed waiting for Pi RPC process: {error}"),
};
FailureReason {
code: ErrorCode::BackendDisconnected,
message,
}
}
async fn force_kill_and_wait(child: &mut Child) {
let _ = child.start_kill();
let _ = child.wait().await;
}
fn command_name(command: &PiRpcCommand) -> &'static str {
match command {
PiRpcCommand::Prompt { .. } => "prompt",
PiRpcCommand::Abort { .. } => "abort",
PiRpcCommand::GetState { .. } => "get_state",
PiRpcCommand::GetMessages { .. } => "get_messages",
PiRpcCommand::GetAvailableModels { .. } => "get_available_models",
PiRpcCommand::SetModel { .. } => "set_model",
PiRpcCommand::SetThinkingLevel { .. } => "set_thinking_level",
PiRpcCommand::GetCommands { .. } => "get_commands",
PiRpcCommand::GetSessionStats { .. } => "get_session_stats",
PiRpcCommand::ExtensionUiResponse { .. } => "extension_ui_response",
}
}
fn reason_from(error: &AgentError) -> FailureReason {
FailureReason {
code: error.code(),
message: error.message().to_owned(),
}
}
fn process_error(message: impl Into<String>) -> AgentError {
AgentError::new(ErrorCode::BackendDisconnected, message)
}
fn invalid_protocol(message: impl Into<String>) -> AgentError {
AgentError::new(ErrorCode::InvalidMessage, message)
}
fn invalid_usage(message: impl Into<String>) -> AgentError {
AgentError::new(ErrorCode::InvalidMessage, message)
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use futures_util::poll;
use tokio::sync::{Semaphore, mpsc, oneshot};
use super::*;
fn event(sequence: u64) -> EventEnvelope {
EventEnvelope {
sequence,
event: PiRpcEvent::AgentStart,
_byte_permit: Arc::new(Semaphore::new(1))
.try_acquire_owned()
.expect("test event byte permit is available"),
}
}
#[tokio::test(flavor = "current_thread")]
async fn fair_retry_yields_after_an_initial_full_queue() {
let (event_tx, mut event_rx) = mpsc::channel::<EventEnvelope>(1);
let (receiver_pending_tx, receiver_pending_rx) = oneshot::channel();
let (first_consumed_tx, first_consumed_rx) = oneshot::channel();
let receiver = tokio::spawn(async move {
let mut first_receive = Box::pin(event_rx.recv());
assert!(poll!(first_receive.as_mut()).is_pending());
receiver_pending_tx
.send(())
.expect("receiver pending signal is observed");
let first = first_receive.await.expect("first event is delivered");
first_consumed_tx
.send(())
.expect("first event consumption is observed");
let second = event_rx.recv().await.expect("retry event is delivered");
(first.sequence, second.sequence)
});
receiver_pending_rx
.await
.expect("receiver has registered its pending receive");
event_tx
.try_send(event(1))
.expect("first event fills the empty queue before the receiver can run");
let mut retry = Box::pin(enqueue_event_with_one_fair_retry(&event_tx, event(2)));
assert!(
poll!(retry.as_mut()).is_pending(),
"the initial full queue must reach the fair yield before retrying"
);
first_consumed_rx
.await
.expect("receiver consumes the first event while retry is yielded");
assert!(
retry.await.is_ok(),
"retry succeeds after the receiver drains one slot"
);
assert_eq!(
receiver.await.expect("receiver task finishes"),
(1, 2),
"both events are delivered in order"
);
}
}