iota/iota-daemon-lib/tests/ipc_server.rs
2026-08-20 17:05:43 +02:00

252 lines
7.4 KiB
Rust

use async_trait::async_trait;
use iota_daemon_lib::{DaemonRuntime, DaemonServices, IpcServer};
use iota_ipc::{
ClientMessage, DaemonMessage, ExitIntent, IpcErrorCode, LocalRequest, PROTOCOL_VERSION,
RequestEnvelope, ResponseResult, read_msg, write_msg,
};
use iota_storage::util::config_util::{self, IotaConfig};
use mtp::codec::CommunicationValue;
use omikron_connector::{OmikronClient, OmikronError};
use std::path::Path;
use std::sync::Arc;
use std::time::Duration;
use tokio::net::UnixStream;
use tokio::sync::{broadcast, watch};
struct ConfigRestore(Arc<IotaConfig>);
impl Drop for ConfigRestore {
fn drop(&mut self) {
config_util::CONFIG.store(self.0.clone());
}
}
fn set_client_limit(limit: usize) -> ConfigRestore {
let previous = config_util::CONFIG.load_full();
let mut config = (*previous).clone();
config.max_ipc_clients = limit;
config_util::CONFIG.store(Arc::new(config));
ConfigRestore(previous)
}
async fn start_server(
path: &Path,
services: Arc<DaemonServices>,
) -> (Arc<DaemonRuntime>, tokio::task::JoinHandle<()>) {
let runtime = Arc::new(DaemonRuntime::new());
let (log_tx, _) = broadcast::channel(32);
let log_buffer = Arc::new(std::sync::Mutex::new(
iota_daemon_lib::log_buffer::LogBuffer::new(32),
));
let (_, state_rx) = watch::channel(runtime.snapshot());
let server = IpcServer::bind(
path.to_owned(),
runtime.clone(),
services,
log_tx,
log_buffer,
state_rx,
)
.await
.expect("IPC server binds");
let task = tokio::spawn(async move {
let _ = server.serve().await;
});
(runtime, task)
}
async fn try_connect_and_await_hello(path: &Path) -> std::io::Result<UnixStream> {
let mut stream = UnixStream::connect(path).await?;
write_msg(
&mut stream,
&ClientMessage::Hello {
supported_versions: vec![PROTOCOL_VERSION],
},
)
.await?;
let message: DaemonMessage = read_msg(&mut stream).await?;
if !matches!(message, DaemonMessage::HelloAck(_)) {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"expected HelloAck",
));
}
Ok(stream)
}
async fn connect_and_await_hello(path: &Path) -> UnixStream {
try_connect_and_await_hello(path)
.await
.expect("IPC connection completes the Hello exchange")
}
#[tokio::test]
async fn active_client_limit_rejects_excess_clients_and_releases_permits() {
let _config = set_client_limit(1);
let directory = tempfile::tempdir().unwrap();
let socket = directory.path().join("ipc.sock");
let (_runtime, server_task) = start_server(&socket, DaemonServices::inactive()).await;
let first = connect_and_await_hello(&socket).await;
let mut rejected = UnixStream::connect(&socket)
.await
.expect("second connection reaches the Unix listener");
let rejected_result = tokio::time::timeout(
Duration::from_secs(2),
read_msg::<_, DaemonMessage>(&mut rejected),
)
.await
.expect("rejected client is closed promptly");
assert!(rejected_result.is_err());
drop(first);
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
let _released = loop {
match try_connect_and_await_hello(&socket).await {
Ok(stream) => break stream,
Err(_error) if tokio::time::Instant::now() < deadline => {
tokio::time::sleep(Duration::from_millis(10)).await;
}
Err(error) => panic!("client permit was not released: {error}"),
}
};
server_task.abort();
let _ = server_task.await;
}
#[tokio::test]
async fn request_with_version_different_from_hello_is_rejected() {
let directory = tempfile::tempdir().unwrap();
let socket = directory.path().join("ipc.sock");
let (_runtime, server_task) = start_server(&socket, DaemonServices::inactive()).await;
let mut stream = connect_and_await_hello(&socket).await;
write_msg(
&mut stream,
&ClientMessage::Request(RequestEnvelope {
request_id: 7,
protocol_version: PROTOCOL_VERSION + 1,
request: LocalRequest::GetStatus,
}),
)
.await
.expect("request sends");
let response = loop {
match read_msg::<_, DaemonMessage>(&mut stream)
.await
.expect("daemon response arrives")
{
DaemonMessage::Response(response) => break response,
_ => continue,
}
};
assert_eq!(response.request_id, 7);
assert!(matches!(
response.result,
ResponseResult::Error(IpcErrorCode::UnsupportedVersion)
));
drop(stream);
server_task.abort();
let _ = server_task.await;
}
struct TestOmikron;
#[async_trait]
impl OmikronClient for TestOmikron {
async fn send_message(&self, _: &CommunicationValue) -> Result<(), OmikronError> {
Ok(())
}
async fn await_response(
&self,
_: &CommunicationValue,
_: Duration,
) -> Result<CommunicationValue, OmikronError> {
Err(OmikronError::Disconnected("test client".into()))
}
async fn reconnect(&self) -> Result<(), OmikronError> {
Ok(())
}
async fn rotate_identity(&self) -> Result<(), OmikronError> {
Ok(())
}
async fn is_connected(&self) -> bool {
true
}
}
fn active_services() -> Arc<DaemonServices> {
Arc::new(DaemonServices {
omikron: Arc::new(TestOmikron),
users: Default::default(),
config: Default::default(),
active: true,
})
}
#[tokio::test]
async fn shutdown_delivers_response_and_lifecycle_event_before_eof() {
let directory = tempfile::tempdir().unwrap();
let socket = directory.path().join("ipc.sock");
let (runtime, server_task) = start_server(&socket, active_services()).await;
let mut stream = connect_and_await_hello(&socket).await;
write_msg(
&mut stream,
&ClientMessage::Request(RequestEnvelope {
request_id: 8,
protocol_version: PROTOCOL_VERSION,
request: LocalRequest::RequestProcessExit {
intent: ExitIntent::Stop,
},
}),
)
.await
.expect("shutdown request sends");
let mut response_seen = false;
let mut lifecycle_seen = false;
for _ in 0..4 {
match tokio::time::timeout(
Duration::from_secs(2),
read_msg::<_, DaemonMessage>(&mut stream),
)
.await
.expect("shutdown message arrives")
.expect("shutdown stream remains readable")
{
DaemonMessage::Response(response) => {
assert_eq!(response.request_id, 8);
assert!(matches!(response.result, ResponseResult::Ok(_)));
response_seen = true;
}
DaemonMessage::LifecycleEvent(iota_ipc::LifecycleEvent::Shutdown { .. }) => {
lifecycle_seen = true;
}
_ => {}
}
if response_seen && lifecycle_seen {
break;
}
}
assert!(response_seen);
assert!(lifecycle_seen);
let eof = tokio::time::timeout(
Duration::from_secs(2),
read_msg::<_, DaemonMessage>(&mut stream),
)
.await
.expect("shutdown connection closes after flush");
assert!(eof.is_err());
assert!(runtime.is_shutting_down());
server_task.abort();
let _ = server_task.await;
}