Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 66 additions & 3 deletions codi-rs/src/orchestrate/ipc/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,12 @@ use super::protocol::{
use super::transport::{self, IpcStream};
use super::super::types::{WorkerConfig, WorkerResult, WorkerStatus, WorkspaceInfo};

const CONNECT_RETRY_ATTEMPTS: usize = 10;
const CONNECT_RETRY_DELAY: Duration = Duration::from_millis(100);
const CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(2);
const HANDSHAKE_POLL_INTERVAL: Duration = Duration::from_millis(10);

/// Error type for IPC client operations.
#[derive(Debug, thiserror::Error)]
pub enum IpcClientError {
Expand Down Expand Up @@ -107,7 +113,33 @@ impl IpcClient {

/// Connect to the commander's endpoint.
pub async fn connect(&mut self) -> Result<(), IpcClientError> {
let stream = transport::connect(&self.socket_path).await?;
let mut last_error: Option<String> = None;
let mut stream = None;

for attempt in 0..CONNECT_RETRY_ATTEMPTS {
match tokio::time::timeout(CONNECT_TIMEOUT, transport::connect(&self.socket_path)).await {
Ok(Ok(conn)) => {
stream = Some(conn);
break;
}
Ok(Err(err)) => {
last_error = Some(err.to_string());
}
Err(_) => {
last_error = Some("connect timeout".to_string());
}
}

if attempt + 1 < CONNECT_RETRY_ATTEMPTS {
tokio::time::sleep(CONNECT_RETRY_DELAY).await;
}
}

let stream = stream.ok_or_else(|| {
IpcClientError::ConnectionFailed(
last_error.unwrap_or_else(|| "failed to connect".to_string())
)
})?;
let (read_half, write_half) = tokio::io::split(stream);

self.writer = Some(write_half);
Expand Down Expand Up @@ -237,7 +269,7 @@ impl IpcClient {
writer.flush().await?;

let ack = self
.wait_for_handshake_ack(Duration::from_secs(2))
.wait_for_handshake_ack(HANDSHAKE_TIMEOUT)
.await;

if let Some(ack) = ack {
Expand Down Expand Up @@ -285,7 +317,7 @@ impl IpcClient {
if let Some(ack) = self.handshake_ack.lock().await.take() {
return ack;
}
tokio::time::sleep(Duration::from_millis(10)).await;
tokio::time::sleep(HANDSHAKE_POLL_INTERVAL).await;
}
})
.await
Expand Down Expand Up @@ -426,4 +458,35 @@ mod tests {
let client = IpcClient::new("/tmp/test.sock", "worker-1");
assert!(!client.is_cancelled().await);
}

#[tokio::test]
async fn test_wait_for_handshake_ack_timeout() {
let client = IpcClient::new("/tmp/test.sock", "worker-1");
let ack = client
.wait_for_handshake_ack(Duration::from_millis(1))
.await;
assert!(ack.is_none());
}

#[tokio::test]
async fn test_wait_for_handshake_ack_returns_value() {
let client = IpcClient::new("/tmp/test.sock", "worker-1");
{
let mut ack = client.handshake_ack.lock().await;
*ack = Some(HandshakeAck {
accepted: true,
auto_approve: Vec::new(),
dangerous_patterns: Vec::new(),
timeout_ms: 123,
reason: None,
});
}

let ack = client
.wait_for_handshake_ack(Duration::from_millis(20))
.await
.expect("ack missing");
assert!(ack.accepted);
assert_eq!(ack.timeout_ms, 123);
}
}
80 changes: 79 additions & 1 deletion codi-rs/src/orchestrate/ipc/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,8 @@ pub use client::IpcClient;

#[cfg(test)]
mod tests {
use super::{CommanderMessage, IpcClient, IpcServer, WorkerMessage};
use super::{CommanderMessage, IpcClient, IpcServer, PermissionResult, WorkerMessage};
use crate::agent::ToolConfirmation;
use crate::orchestrate::types::{WorkerConfig, WorkspaceInfo};
use std::path::PathBuf;
use std::sync::Arc;
Expand Down Expand Up @@ -123,4 +124,81 @@ mod tests {
accept_task.await.expect("accept task failed");
ack_task.await.expect("ack task failed");
}

#[cfg(windows)]
#[tokio::test]
async fn test_named_pipe_permission_roundtrip() {
let pipe_name = format!(r"\\.\pipe\codi-ipc-permission-{}", uuid::Uuid::new_v4());
let socket_path = PathBuf::from(pipe_name);

let mut server = IpcServer::new(&socket_path);
server.start().await.expect("server start failed");

let mut rx = server.take_receiver().expect("receiver already taken");
let server = Arc::new(server);

let accept_server = Arc::clone(&server);
let accept_task = tokio::spawn(async move {
accept_server.accept().await.expect("accept failed")
});

let ack_server = Arc::clone(&server);
let server_task = tokio::spawn(async move {
let (worker_id, msg) = rx.recv().await.expect("handshake missing");
assert_eq!(worker_id, "worker-1");
assert!(matches!(msg, WorkerMessage::Handshake { .. }));

let ack = CommanderMessage::handshake_ack(
true,
Vec::new(),
Vec::new(),
5_000,
);
ack_server
.send(&worker_id, &ack)
.await
.expect("ack send failed");

let (worker_id, msg) = rx.recv().await.expect("permission missing");
if let WorkerMessage::PermissionRequest { request_id, .. } = msg {
ack_server
.send(&worker_id, &CommanderMessage::approve(request_id))
.await
.expect("permission approve failed");
} else {
panic!("expected permission request");
}
});

let mut client = IpcClient::new(&socket_path, "worker-1");
client.connect().await.expect("client connect failed");

let workspace = WorkspaceInfo::GitWorktree {
path: PathBuf::from("."),
branch: "feat/test".to_string(),
base_branch: "main".to_string(),
};
let config = WorkerConfig::new("worker-1", "feat/test", "task");

client
.handshake(&config, &workspace)
.await
.expect("handshake failed");

let confirmation = ToolConfirmation {
tool_name: "read_file".to_string(),
input: serde_json::json!({ "path": "README.md" }),
is_dangerous: false,
danger_reason: None,
};

let result = client
.request_permission(&confirmation)
.await
.expect("permission request failed");
assert_eq!(result, PermissionResult::Approve);

accept_task.await.expect("accept task failed");
server_task.await.expect("server task failed");
}
}