use std::path::Path;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use dashmap::DashMap;
use serde::Deserialize;
use tokio::net::UnixStream;
use tokio::sync::{mpsc, oneshot};
use tokio::task::JoinHandle;
use tracing::{debug, error, trace, warn};
use super::protocol::{
read_frame, write_frame, Format, MAX_FRAME_SIZE, MSG_TYPE_DISCOVER, MSG_TYPE_ERROR,
MSG_TYPE_HEARTBEAT, MSG_TYPE_PUSH, MSG_TYPE_REQUEST, MSG_TYPE_RESPONSE, MSG_TYPE_STREAM,
MSG_TYPE_SUBSCRIBE, MSG_TYPE_UNSUBSCRIBE,
};
use super::types::{
IpcDiscoverRequest, IpcDiscoverResponse, IpcEnvelope, IpcError, IpcPushNotification,
IpcResponse, IpcStreamFrame, IpcSubscribeRequest, IpcSubscriptionResponse,
IpcUnsubscribeRequest,
};
const DEFAULT_WRITER_CHANNEL_CAPACITY: usize = 64;
const DEFAULT_PUSH_CHANNEL_CAPACITY: usize = 256;
const DEFAULT_STREAM_CHANNEL_CAPACITY: usize = 64;
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
type RawResponse = Result<(Format, Vec<u8>), IpcError>;
type PendingRequests = DashMap<String, oneshot::Sender<RawResponse>>;
type ActiveStreams = DashMap<String, mpsc::Sender<IpcStreamFrame>>;
enum ClientWriteCommand {
Envelope { envelope: IpcEnvelope, format: Format },
Request {
envelope: IpcEnvelope,
format: Format,
reply_tx: oneshot::Sender<RawResponse>,
},
StreamRequest { envelope: IpcEnvelope, format: Format },
Subscribe {
request: IpcSubscribeRequest,
format: Format,
reply_tx: oneshot::Sender<RawResponse>,
},
Unsubscribe {
request: IpcUnsubscribeRequest,
format: Format,
reply_tx: oneshot::Sender<RawResponse>,
},
Discover {
request: IpcDiscoverRequest,
format: Format,
reply_tx: oneshot::Sender<RawResponse>,
},
Shutdown,
}
#[derive(Deserialize)]
struct CorrelationPeek {
correlation_id: String,
}
#[derive(Debug, Clone)]
pub struct IpcClientConfig {
pub format: Format,
pub writer_channel_capacity: usize,
pub push_channel_capacity: usize,
pub default_timeout: Duration,
pub max_frame_size: usize,
}
impl Default for IpcClientConfig {
fn default() -> Self {
Self {
format: Format::default(),
writer_channel_capacity: DEFAULT_WRITER_CHANNEL_CAPACITY,
push_channel_capacity: DEFAULT_PUSH_CHANNEL_CAPACITY,
default_timeout: DEFAULT_TIMEOUT,
max_frame_size: MAX_FRAME_SIZE,
}
}
}
pub struct IpcClient {
writer_tx: mpsc::Sender<ClientWriteCommand>,
push_rx: std::sync::Mutex<Option<mpsc::Receiver<IpcPushNotification>>>,
active_streams: std::sync::Arc<ActiveStreams>,
format: Format,
default_timeout: Duration,
reader_handle: JoinHandle<()>,
writer_handle: std::sync::Mutex<Option<JoinHandle<()>>>,
shutdown: AtomicBool,
}
impl IpcClient {
pub async fn connect(path: impl AsRef<Path>) -> Result<Self, IpcError> {
Self::connect_with_config(path, IpcClientConfig::default()).await
}
pub async fn connect_with_config(
path: impl AsRef<Path>,
config: IpcClientConfig,
) -> Result<Self, IpcError> {
let path = path.as_ref();
debug!(path = %path.display(), "IPC client connecting");
let stream = UnixStream::connect(path)
.await
.map_err(|e| IpcError::IoError(format!("Failed to connect to {}: {e}", path.display())))?;
let (reader, writer) = stream.into_split();
let (writer_tx, writer_rx) =
mpsc::channel::<ClientWriteCommand>(config.writer_channel_capacity);
let (push_tx, push_rx) =
mpsc::channel::<IpcPushNotification>(config.push_channel_capacity);
let pending_requests: std::sync::Arc<PendingRequests> =
std::sync::Arc::new(DashMap::new());
let pending_for_writer = std::sync::Arc::clone(&pending_requests);
let active_streams: std::sync::Arc<ActiveStreams> = std::sync::Arc::new(DashMap::new());
let streams_for_writer = std::sync::Arc::clone(&active_streams);
let streams_for_reader = std::sync::Arc::clone(&active_streams);
let writer_handle = tokio::spawn(async move {
run_client_writer_task(writer, writer_rx, pending_for_writer, streams_for_writer)
.await;
});
let max_frame_size = config.max_frame_size;
let reader_handle = tokio::spawn(async move {
run_client_reader_task(
reader,
pending_requests,
streams_for_reader,
push_tx,
max_frame_size,
)
.await;
});
debug!(path = %path.display(), "IPC client connected");
Ok(Self {
writer_tx,
push_rx: std::sync::Mutex::new(Some(push_rx)),
active_streams,
format: config.format,
default_timeout: config.default_timeout,
reader_handle,
writer_handle: std::sync::Mutex::new(Some(writer_handle)),
shutdown: AtomicBool::new(false),
})
}
pub async fn send(&self, envelope: IpcEnvelope) -> Result<(), IpcError> {
self.writer_tx
.send(ClientWriteCommand::Envelope {
envelope,
format: self.format,
})
.await
.map_err(|_| IpcError::ConnectionClosed)
}
pub async fn request(&self, envelope: IpcEnvelope) -> Result<IpcResponse, IpcError> {
self.request_with_timeout(envelope, self.default_timeout)
.await
}
pub async fn request_with_timeout(
&self,
envelope: IpcEnvelope,
timeout_duration: Duration,
) -> Result<IpcResponse, IpcError> {
let (reply_tx, reply_rx) = oneshot::channel();
self.writer_tx
.send(ClientWriteCommand::Request {
envelope,
format: self.format,
reply_tx,
})
.await
.map_err(|_| IpcError::ConnectionClosed)?;
let (format, bytes) = tokio::time::timeout(timeout_duration, reply_rx)
.await
.map_err(|_| IpcError::Timeout)?
.map_err(|_| IpcError::ConnectionClosed)??;
format.deserialize(&bytes)
}
pub async fn request_stream(
&self,
envelope: IpcEnvelope,
) -> Result<mpsc::Receiver<IpcStreamFrame>, IpcError> {
self.request_stream_with_timeout(envelope, self.default_timeout)
.await
}
pub async fn request_stream_with_timeout(
&self,
mut envelope: IpcEnvelope,
frame_timeout: Duration,
) -> Result<mpsc::Receiver<IpcStreamFrame>, IpcError> {
envelope.expects_stream = true;
envelope.expects_reply = false;
let correlation_id = envelope.correlation_id.clone();
let (frame_tx, frame_rx) = mpsc::channel(DEFAULT_STREAM_CHANNEL_CAPACITY);
let (out_tx, out_rx) = mpsc::channel(DEFAULT_STREAM_CHANNEL_CAPACITY);
self.active_streams.insert(correlation_id.clone(), frame_tx);
if self
.writer_tx
.send(ClientWriteCommand::StreamRequest {
envelope,
format: self.format,
})
.await
.is_err()
{
self.active_streams.remove(&correlation_id);
return Err(IpcError::ConnectionClosed);
}
tokio::spawn(run_stream_forwarder(
correlation_id,
frame_rx,
out_tx,
frame_timeout,
std::sync::Arc::clone(&self.active_streams),
));
Ok(out_rx)
}
pub async fn subscribe(
&self,
message_types: Vec<String>,
) -> Result<IpcSubscriptionResponse, IpcError> {
let request = IpcSubscribeRequest::new(message_types);
let (reply_tx, reply_rx) = oneshot::channel();
self.writer_tx
.send(ClientWriteCommand::Subscribe {
request,
format: self.format,
reply_tx,
})
.await
.map_err(|_| IpcError::ConnectionClosed)?;
let (format, bytes) = tokio::time::timeout(self.default_timeout, reply_rx)
.await
.map_err(|_| IpcError::Timeout)?
.map_err(|_| IpcError::ConnectionClosed)??;
format.deserialize(&bytes)
}
pub async fn unsubscribe(
&self,
message_types: Vec<String>,
) -> Result<IpcSubscriptionResponse, IpcError> {
let request = if message_types.is_empty() {
IpcUnsubscribeRequest::unsubscribe_all()
} else {
IpcUnsubscribeRequest::new(message_types)
};
let (reply_tx, reply_rx) = oneshot::channel();
self.writer_tx
.send(ClientWriteCommand::Unsubscribe {
request,
format: self.format,
reply_tx,
})
.await
.map_err(|_| IpcError::ConnectionClosed)?;
let (format, bytes) = tokio::time::timeout(self.default_timeout, reply_rx)
.await
.map_err(|_| IpcError::Timeout)?
.map_err(|_| IpcError::ConnectionClosed)??;
format.deserialize(&bytes)
}
pub async fn discover(&self) -> Result<IpcDiscoverResponse, IpcError> {
let request = IpcDiscoverRequest::new();
let (reply_tx, reply_rx) = oneshot::channel();
self.writer_tx
.send(ClientWriteCommand::Discover {
request,
format: self.format,
reply_tx,
})
.await
.map_err(|_| IpcError::ConnectionClosed)?;
let (format, bytes) = tokio::time::timeout(self.default_timeout, reply_rx)
.await
.map_err(|_| IpcError::Timeout)?
.map_err(|_| IpcError::ConnectionClosed)??;
format.deserialize(&bytes)
}
pub fn take_push_receiver(&self) -> Option<mpsc::Receiver<IpcPushNotification>> {
self.push_rx
.lock()
.ok()
.and_then(|mut guard| guard.take())
}
#[must_use]
pub const fn format(&self) -> Format {
self.format
}
#[must_use]
pub fn is_connected(&self) -> bool {
if self.shutdown.load(Ordering::Relaxed) {
return false;
}
self.writer_handle
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.as_ref()
.is_some_and(|h| !h.is_finished())
}
pub async fn disconnect(&self) -> Result<(), IpcError> {
if self.shutdown.swap(true, Ordering::SeqCst) {
return Ok(()); }
let request = IpcUnsubscribeRequest::unsubscribe_all();
let payload = self.format.serialize(&request)?;
let _ = self
.writer_tx
.send(ClientWriteCommand::Envelope {
envelope: IpcEnvelope::new(
"__unsubscribe__",
"IpcUnsubscribeRequest",
serde_json::Value::Null,
),
format: self.format,
})
.await;
drop(payload);
let _ = self
.writer_tx
.send(ClientWriteCommand::Shutdown)
.await;
let writer_handle = self
.writer_handle
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
if let Some(handle) = writer_handle {
let _ = handle.await;
}
self.reader_handle.abort();
debug!("IPC client disconnected");
Ok(())
}
}
impl Drop for IpcClient {
fn drop(&mut self) {
self.reader_handle.abort();
let writer_handle = self
.writer_handle
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
if let Some(handle) = writer_handle {
handle.abort();
}
self.active_streams.clear();
}
}
async fn write_correlated_frame<T: serde::Serialize + Sync>(
writer: &mut tokio::net::unix::OwnedWriteHalf,
pending_requests: &PendingRequests,
correlation_id: String,
reply_tx: oneshot::Sender<RawResponse>,
msg_type: u8,
format: Format,
value: &T,
) -> Option<Result<(), IpcError>> {
pending_requests.insert(correlation_id.clone(), reply_tx);
let payload = match format.serialize(value) {
Ok(p) => p,
Err(e) => {
if let Some((_, tx)) = pending_requests.remove(&correlation_id) {
let _ = tx.send(Err(e.clone()));
}
error!(error = %e, "Failed to serialize IPC frame");
return None;
}
};
let result = write_frame(writer, msg_type, format, &payload).await;
if let Err(ref e) = result {
if let Some((_, tx)) = pending_requests.remove(&correlation_id) {
let _ = tx.send(Err(e.clone()));
}
}
Some(result)
}
async fn run_client_writer_task(
mut writer: tokio::net::unix::OwnedWriteHalf,
mut receiver: mpsc::Receiver<ClientWriteCommand>,
pending_requests: std::sync::Arc<PendingRequests>,
active_streams: std::sync::Arc<ActiveStreams>,
) {
trace!("IPC client writer task started");
while let Some(cmd) = receiver.recv().await {
let result = match cmd {
ClientWriteCommand::Envelope { envelope, format } => {
let payload = match format.serialize(&envelope) {
Ok(p) => p,
Err(e) => {
error!(error = %e, "Failed to serialize envelope");
continue;
}
};
Some(write_frame(&mut writer, MSG_TYPE_REQUEST, format, &payload).await)
}
ClientWriteCommand::Request { envelope, format, reply_tx } => {
let cid = envelope.correlation_id.clone();
write_correlated_frame(
&mut writer, &pending_requests, cid, reply_tx,
MSG_TYPE_REQUEST, format, &envelope,
).await
}
ClientWriteCommand::StreamRequest { envelope, format } => {
let payload = match format.serialize(&envelope) {
Ok(p) => p,
Err(e) => {
error!(error = %e, "Failed to serialize stream request");
active_streams.remove(&envelope.correlation_id);
continue;
}
};
let result = write_frame(&mut writer, MSG_TYPE_REQUEST, format, &payload).await;
if result.is_err() {
active_streams.remove(&envelope.correlation_id);
}
Some(result)
}
ClientWriteCommand::Subscribe { request, format, reply_tx } => {
let cid = request.correlation_id.clone();
write_correlated_frame(
&mut writer, &pending_requests, cid, reply_tx,
MSG_TYPE_SUBSCRIBE, format, &request,
).await
}
ClientWriteCommand::Unsubscribe { request, format, reply_tx } => {
let cid = request.correlation_id.clone();
write_correlated_frame(
&mut writer, &pending_requests, cid, reply_tx,
MSG_TYPE_UNSUBSCRIBE, format, &request,
).await
}
ClientWriteCommand::Discover { request, format, reply_tx } => {
let cid = request.correlation_id.clone();
write_correlated_frame(
&mut writer, &pending_requests, cid, reply_tx,
MSG_TYPE_DISCOVER, format, &request,
).await
}
ClientWriteCommand::Shutdown => {
trace!("IPC client writer received shutdown command");
break;
}
};
if let Some(Err(e)) = result {
error!(error = %e, "IPC client writer error, closing connection");
break;
}
}
pending_requests.clear();
active_streams.clear();
trace!("IPC client writer task finished");
}
async fn run_client_reader_task(
mut reader: tokio::net::unix::OwnedReadHalf,
pending_requests: std::sync::Arc<PendingRequests>,
active_streams: std::sync::Arc<ActiveStreams>,
push_tx: mpsc::Sender<IpcPushNotification>,
max_frame_size: usize,
) {
trace!("IPC client reader task started");
loop {
match read_frame(&mut reader, max_frame_size).await {
Ok((msg_type, format, payload)) => match msg_type {
MSG_TYPE_RESPONSE | MSG_TYPE_ERROR => {
handle_response_frame(&pending_requests, &active_streams, format, payload)
.await;
}
MSG_TYPE_STREAM => {
handle_stream_frame(&active_streams, format, &payload).await;
}
MSG_TYPE_PUSH => {
handle_push_frame(&push_tx, format, &payload);
}
MSG_TYPE_HEARTBEAT => {
trace!("IPC client received heartbeat");
}
_ => {
warn!(msg_type, "IPC client received unknown message type");
}
},
Err(IpcError::ConnectionClosed) => {
debug!("IPC client connection closed by server");
break;
}
Err(e) => {
error!(error = %e, "IPC client reader error");
break;
}
}
}
for entry in pending_requests.iter() {
let correlation_id = entry.key().clone();
if let Some((_, tx)) = pending_requests.remove(&correlation_id) {
let _ = tx.send(Err(IpcError::ConnectionClosed));
}
}
active_streams.clear();
trace!("IPC client reader task finished");
}
async fn handle_response_frame(
pending_requests: &PendingRequests,
active_streams: &ActiveStreams,
format: Format,
payload: Vec<u8>,
) {
let correlation_id = match format.deserialize::<CorrelationPeek>(&payload) {
Ok(peek) => peek.correlation_id,
Err(e) => {
warn!(error = %e, "Failed to peek correlation_id from response");
return;
}
};
if let Some((_, reply_tx)) = pending_requests.remove(&correlation_id) {
let _ = reply_tx.send(Ok((format, payload)));
return;
}
if let Some(frame_tx) = active_streams
.get(&correlation_id)
.map(|entry| entry.value().clone())
{
debug!(correlation_id, "Server rejected stream request, terminating stream");
let frame = synthesize_stream_termination(&correlation_id, format, &payload);
let _ = frame_tx.send(frame).await;
active_streams.remove(&correlation_id);
return;
}
trace!(correlation_id, "Draining unclaimed response");
}
fn synthesize_stream_termination(
correlation_id: &str,
format: Format,
payload: &[u8],
) -> IpcStreamFrame {
match format.deserialize::<IpcResponse>(payload) {
Ok(response) if response.success => IpcStreamFrame {
correlation_id: correlation_id.to_string(),
sequence: 0,
is_final: true,
error: None,
error_code: None,
payload: response.payload,
},
Ok(response) => IpcStreamFrame {
correlation_id: correlation_id.to_string(),
sequence: 0,
is_final: true,
error: Some(response.error.unwrap_or_else(|| {
"Server rejected the stream request".to_string()
})),
error_code: response.error_code,
payload: response.payload,
},
Err(e) => IpcStreamFrame::error(
correlation_id,
0,
format!("Server rejected the stream request with an undecodable response: {e}"),
),
}
}
async fn handle_stream_frame(active_streams: &ActiveStreams, format: Format, payload: &[u8]) {
let frame = match format.deserialize::<IpcStreamFrame>(payload) {
Ok(frame) => frame,
Err(e) => {
warn!(error = %e, "Failed to deserialize stream frame");
return;
}
};
let correlation_id = frame.correlation_id.clone();
let is_final = frame.is_final;
let Some(frame_tx) = active_streams
.get(&correlation_id)
.map(|entry| entry.value().clone())
else {
warn!(correlation_id, "Received stream frame for unknown correlation_id, dropping");
return;
};
if frame_tx.send(frame).await.is_err() {
debug!(correlation_id, "Stream frame channel closed, removing stream");
active_streams.remove(&correlation_id);
return;
}
if is_final {
active_streams.remove(&correlation_id);
}
}
async fn run_stream_forwarder(
correlation_id: String,
mut frame_rx: mpsc::Receiver<IpcStreamFrame>,
out_tx: mpsc::Sender<IpcStreamFrame>,
frame_timeout: Duration,
active_streams: std::sync::Arc<ActiveStreams>,
) {
trace!(correlation_id, "Stream forwarder started");
loop {
match tokio::time::timeout(frame_timeout, frame_rx.recv()).await {
Ok(Some(frame)) => {
let is_final = frame.is_final;
if out_tx.send(frame).await.is_err() {
debug!(correlation_id, "Stream receiver dropped, cancelling stream");
break;
}
if is_final {
trace!(correlation_id, "Stream completed with final frame");
break;
}
}
Ok(None) => {
debug!(correlation_id, "Stream terminated before final frame");
break;
}
Err(_) => {
warn!(
correlation_id,
timeout = ?frame_timeout,
"Stream timed out waiting for next frame"
);
break;
}
}
}
active_streams.remove(&correlation_id);
trace!(correlation_id, "Stream forwarder finished");
}
fn handle_push_frame(
push_tx: &mpsc::Sender<IpcPushNotification>,
format: Format,
payload: &[u8],
) {
match format.deserialize::<IpcPushNotification>(payload) {
Ok(notification) => {
if let Err(e) = push_tx.try_send(notification) {
match e {
mpsc::error::TrySendError::Full(_) => {
warn!("Push notification channel full, dropping notification");
}
mpsc::error::TrySendError::Closed(_) => {
debug!("Push notification channel closed");
}
}
}
}
Err(e) => {
warn!(error = %e, "Failed to deserialize push notification");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn config_defaults() {
let config = IpcClientConfig::default();
assert_eq!(config.writer_channel_capacity, 64);
assert_eq!(config.push_channel_capacity, 256);
assert_eq!(config.default_timeout, Duration::from_secs(30));
assert_eq!(config.max_frame_size, MAX_FRAME_SIZE);
}
#[tokio::test]
async fn connect_to_socket() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
stream
});
let client = IpcClient::connect(&socket_path).await;
assert!(client.is_ok());
let client = client.expect("client should connect");
assert!(client.is_connected());
let _server_stream = accept_handle.await.expect("accept failed");
drop(client);
}
#[tokio::test]
async fn connect_fails_for_missing_socket() {
let result = IpcClient::connect("/tmp/nonexistent_test_socket_12345.sock").await;
assert!(result.is_err());
}
#[tokio::test]
async fn send_fire_and_forget() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let result = read_frame(&mut reader, MAX_FRAME_SIZE).await;
assert!(result.is_ok());
let (msg_type, _, payload) = result.expect("should read frame");
assert_eq!(msg_type, MSG_TYPE_REQUEST);
let envelope: IpcEnvelope =
serde_json::from_slice(&payload).expect("should deserialize");
assert_eq!(envelope.target, "test_actor");
let response = IpcResponse::success(&envelope.correlation_id, None);
let resp_payload =
serde_json::to_vec(&response).expect("should serialize response");
write_frame(
&mut writer,
MSG_TYPE_RESPONSE,
Format::Json,
&resp_payload,
)
.await
.expect("should write response");
tokio::time::sleep(Duration::from_millis(100)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new(
"test_actor",
"TestMessage",
serde_json::json!({"key": "value"}),
);
let result = client.send(envelope).await;
assert!(result.is_ok());
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn request_response_correlation() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let (msg_type, _, payload) = read_frame(&mut reader, MAX_FRAME_SIZE)
.await
.expect("should read frame");
assert_eq!(msg_type, MSG_TYPE_REQUEST);
let envelope: IpcEnvelope =
serde_json::from_slice(&payload).expect("should deserialize");
let response = IpcResponse::success(
&envelope.correlation_id,
Some(serde_json::json!({"result": 42})),
);
let resp_payload =
serde_json::to_vec(&response).expect("should serialize response");
write_frame(
&mut writer,
MSG_TYPE_RESPONSE,
Format::Json,
&resp_payload,
)
.await
.expect("should write response");
tokio::time::sleep(Duration::from_millis(100)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_request(
"test_actor",
"TestQuery",
serde_json::json!({"query": "test"}),
);
let response = client.request(envelope).await;
assert!(response.is_ok());
let response = response.expect("should get response");
assert!(response.success);
assert_eq!(
response.payload,
Some(serde_json::json!({"result": 42}))
);
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn subscribe_and_receive_push() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let (msg_type, _, payload) = read_frame(&mut reader, MAX_FRAME_SIZE)
.await
.expect("should read frame");
assert_eq!(msg_type, MSG_TYPE_SUBSCRIBE);
let sub_request: IpcSubscribeRequest =
serde_json::from_slice(&payload).expect("should deserialize");
let response = IpcSubscriptionResponse::success(
&sub_request.correlation_id,
sub_request.message_types.clone(),
);
let resp_payload =
serde_json::to_vec(&response).expect("should serialize");
write_frame(
&mut writer,
MSG_TYPE_RESPONSE,
Format::Json,
&resp_payload,
)
.await
.expect("should write response");
let notification = IpcPushNotification::new(
"TestEvent",
Some("test_actor".to_string()),
serde_json::json!({"data": "hello"}),
);
let push_payload =
serde_json::to_vec(¬ification).expect("should serialize");
write_frame(
&mut writer,
MSG_TYPE_PUSH,
Format::Json,
&push_payload,
)
.await
.expect("should write push");
tokio::time::sleep(Duration::from_millis(200)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let sub_response = client
.subscribe(vec!["TestEvent".to_string()])
.await
.expect("should subscribe");
assert!(sub_response.success);
let mut push_rx = client
.take_push_receiver()
.expect("should take push receiver");
let notification = tokio::time::timeout(Duration::from_secs(2), push_rx.recv())
.await
.expect("should not timeout")
.expect("should receive notification");
assert_eq!(notification.message_type, "TestEvent");
assert_eq!(
notification.payload,
serde_json::json!({"data": "hello"})
);
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn request_stream_receives_frames_until_final() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let (msg_type, _, payload) = read_frame(&mut reader, MAX_FRAME_SIZE)
.await
.expect("should read frame");
assert_eq!(msg_type, MSG_TYPE_REQUEST);
let envelope: IpcEnvelope =
serde_json::from_slice(&payload).expect("should deserialize");
assert!(envelope.expects_stream);
for sequence in 0..3_u32 {
let frame = IpcStreamFrame::data(
&envelope.correlation_id,
sequence,
serde_json::json!({ "n": sequence }),
);
let frame_payload =
serde_json::to_vec(&frame).expect("should serialize frame");
write_frame(&mut writer, MSG_TYPE_STREAM, Format::Json, &frame_payload)
.await
.expect("should write frame");
}
let final_frame = IpcStreamFrame::final_frame(&envelope.correlation_id, 3, None);
let final_payload =
serde_json::to_vec(&final_frame).expect("should serialize final frame");
write_frame(&mut writer, MSG_TYPE_STREAM, Format::Json, &final_payload)
.await
.expect("should write final frame");
tokio::time::sleep(Duration::from_millis(200)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_stream_request(
"test_actor",
"TestStream",
serde_json::json!({}),
);
let mut stream_rx = client
.request_stream(envelope)
.await
.expect("should start stream");
let mut frames = Vec::new();
while let Some(frame) =
tokio::time::timeout(Duration::from_secs(2), stream_rx.recv())
.await
.expect("should not timeout")
{
frames.push(frame);
}
assert_eq!(frames.len(), 4);
for (expected_sequence, frame) in (0_u32..).zip(frames.iter()) {
assert_eq!(frame.sequence, expected_sequence);
}
assert!(frames.last().expect("frames should not be empty").is_final);
assert!(frames[..3].iter().all(|f| !f.is_final));
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(client.active_streams.is_empty());
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn request_stream_times_out_without_frames() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, _writer) = stream.into_split();
let _ = read_frame(&mut reader, MAX_FRAME_SIZE).await;
tokio::time::sleep(Duration::from_millis(500)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_stream_request(
"test_actor",
"TestStream",
serde_json::json!({}),
);
let mut stream_rx = client
.request_stream_with_timeout(envelope, Duration::from_millis(100))
.await
.expect("should start stream");
let result = tokio::time::timeout(Duration::from_secs(2), stream_rx.recv())
.await
.expect("stream should terminate before the outer timeout");
assert!(result.is_none());
assert!(client.active_streams.is_empty());
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn request_stream_terminates_on_connection_close() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let (_, _, payload) = read_frame(&mut reader, MAX_FRAME_SIZE)
.await
.expect("should read frame");
let envelope: IpcEnvelope =
serde_json::from_slice(&payload).expect("should deserialize");
for sequence in 0..2_u32 {
let frame = IpcStreamFrame::data(
&envelope.correlation_id,
sequence,
serde_json::json!({ "n": sequence }),
);
let frame_payload =
serde_json::to_vec(&frame).expect("should serialize frame");
write_frame(&mut writer, MSG_TYPE_STREAM, Format::Json, &frame_payload)
.await
.expect("should write frame");
}
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_stream_request(
"test_actor",
"TestStream",
serde_json::json!({}),
);
let mut stream_rx = client
.request_stream(envelope)
.await
.expect("should start stream");
let mut frames = Vec::new();
while let Some(frame) =
tokio::time::timeout(Duration::from_secs(2), stream_rx.recv())
.await
.expect("stream should terminate, not hang")
{
frames.push(frame);
}
assert_eq!(frames.len(), 2);
assert!(frames.iter().all(|f| !f.is_final));
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn request_stream_receiver_drop_cleans_up() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let (_, _, payload) = read_frame(&mut reader, MAX_FRAME_SIZE)
.await
.expect("should read frame");
let envelope: IpcEnvelope =
serde_json::from_slice(&payload).expect("should deserialize");
tokio::time::sleep(Duration::from_millis(100)).await;
let frame = IpcStreamFrame::data(
&envelope.correlation_id,
0,
serde_json::json!({ "n": 0 }),
);
let frame_payload =
serde_json::to_vec(&frame).expect("should serialize frame");
write_frame(&mut writer, MSG_TYPE_STREAM, Format::Json, &frame_payload)
.await
.expect("should write frame");
tokio::time::sleep(Duration::from_millis(200)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_stream_request(
"test_actor",
"TestStream",
serde_json::json!({}),
);
let stream_rx = client
.request_stream(envelope)
.await
.expect("should start stream");
assert_eq!(client.active_streams.len(), 1);
drop(stream_rx);
tokio::time::sleep(Duration::from_millis(300)).await;
assert!(client.active_streams.is_empty());
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn take_push_receiver_once() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
tokio::time::sleep(Duration::from_millis(200)).await;
drop(stream);
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
assert!(client.take_push_receiver().is_some());
assert!(client.take_push_receiver().is_none());
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn request_stream_slow_consumer_receives_all_frames() {
const FRAME_COUNT: u32 = 300;
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let (_, _, payload) = read_frame(&mut reader, MAX_FRAME_SIZE)
.await
.expect("should read frame");
let envelope: IpcEnvelope =
serde_json::from_slice(&payload).expect("should deserialize");
for sequence in 0..FRAME_COUNT - 1 {
let frame = IpcStreamFrame::data(
&envelope.correlation_id,
sequence,
serde_json::json!({ "n": sequence }),
);
let frame_payload =
serde_json::to_vec(&frame).expect("should serialize frame");
write_frame(&mut writer, MSG_TYPE_STREAM, Format::Json, &frame_payload)
.await
.expect("should write frame");
}
let final_frame =
IpcStreamFrame::final_frame(&envelope.correlation_id, FRAME_COUNT - 1, None);
let final_payload =
serde_json::to_vec(&final_frame).expect("should serialize final frame");
write_frame(&mut writer, MSG_TYPE_STREAM, Format::Json, &final_payload)
.await
.expect("should write final frame");
tokio::time::sleep(Duration::from_secs(2)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_stream_request(
"test_actor",
"TestStream",
serde_json::json!({}),
);
let mut stream_rx = client
.request_stream(envelope)
.await
.expect("should start stream");
tokio::time::sleep(Duration::from_millis(300)).await;
let mut frames = Vec::new();
while let Some(frame) =
tokio::time::timeout(Duration::from_secs(5), stream_rx.recv())
.await
.expect("stream should terminate, not hang")
{
frames.push(frame);
}
assert_eq!(frames.len(), FRAME_COUNT as usize);
for (expected_sequence, frame) in (0_u32..).zip(frames.iter()) {
assert_eq!(frame.sequence, expected_sequence);
}
assert!(frames.last().expect("frames should not be empty").is_final);
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn request_stream_terminated_by_rate_limit_rejection() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, mut writer) = stream.into_split();
let (_, _, payload) = read_frame(&mut reader, MAX_FRAME_SIZE)
.await
.expect("should read frame");
let envelope: IpcEnvelope =
serde_json::from_slice(&payload).expect("should deserialize");
let response = IpcResponse::error(
&envelope.correlation_id,
&IpcError::RateLimited { retry_after_ms: 100 },
);
let resp_payload =
serde_json::to_vec(&response).expect("should serialize response");
write_frame(&mut writer, MSG_TYPE_ERROR, Format::Json, &resp_payload)
.await
.expect("should write response");
tokio::time::sleep(Duration::from_millis(200)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_stream_request(
"test_actor",
"TestStream",
serde_json::json!({}),
);
let mut stream_rx = client
.request_stream(envelope)
.await
.expect("should start stream");
let frame = tokio::time::timeout(Duration::from_secs(2), stream_rx.recv())
.await
.expect("stream should terminate, not hang")
.expect("should receive the rejection frame");
assert!(frame.is_final);
assert_eq!(frame.error_code.as_deref(), Some("RATE_LIMITED"));
assert!(frame.error.is_some());
let next = tokio::time::timeout(Duration::from_secs(2), stream_rx.recv())
.await
.expect("stream should terminate, not hang");
assert!(next.is_none());
assert!(client.active_streams.is_empty());
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn drop_client_closes_stream_channel() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
let (mut reader, _writer) = stream.into_split();
let _ = read_frame(&mut reader, MAX_FRAME_SIZE).await;
tokio::time::sleep(Duration::from_millis(500)).await;
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
let envelope = IpcEnvelope::new_stream_request(
"test_actor",
"TestStream",
serde_json::json!({}),
);
let mut stream_rx = client
.request_stream(envelope)
.await
.expect("should start stream");
drop(client);
let result = tokio::time::timeout(Duration::from_secs(2), stream_rx.recv())
.await
.expect("stream should close promptly after drop, not hang");
assert!(result.is_none());
accept_handle.await.expect("server task failed");
}
#[tokio::test]
async fn graceful_disconnect() {
let dir = tempfile::tempdir().expect("failed to create temp dir");
let socket_path = dir.path().join("test.sock");
let listener = tokio::net::UnixListener::bind(&socket_path)
.expect("failed to bind socket");
let accept_handle = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("failed to accept");
tokio::time::sleep(Duration::from_millis(500)).await;
drop(stream);
});
let client = IpcClient::connect(&socket_path)
.await
.expect("should connect");
assert!(client.is_connected());
let result = client.disconnect().await;
assert!(result.is_ok());
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!client.is_connected());
let result = client.disconnect().await;
assert!(result.is_ok());
accept_handle.await.expect("server task failed");
}
}