use futures_util::{SinkExt, StreamExt};
use prost::Message;
use std::sync::{Arc, Condvar, Mutex};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use super::{
serve_broker_session_endpoint, serve_broker_session_endpoint_concurrently,
serve_broker_session_socket, ENVELOPE_VERSION,
};
use crate::broker::protocol::{
encode_framed, hello_reply::Result as HelloReplyResult, Frame, HelloReply, Negotiated,
CONTROL_PAYLOAD_PROTOCOL,
};
use crate::broker::protocol_v2::{session_frame, SessionFrame, SessionStart};
use crate::broker::server::connection::{HelloResponder, PeerCredentialPolicy};
use crate::broker::server::hello_handler::PeerIdentity;
use crate::daemon::compile_session::session_framed;
use crate::daemon::session_endpoint::serve_session_endpoint;
use crate::platform::ipc::{AsyncListener, AsyncStream as IpcAsyncStream, Endpoint as IpcEndpoint};
struct NegotiateToBackend(String);
impl HelloResponder for NegotiateToBackend {
fn handle_frame(&self, _frame: Frame, _peer: PeerIdentity) -> HelloReply {
HelloReply {
result: Some(HelloReplyResult::Negotiated(Negotiated {
backend_pipe: self.0.clone(),
..Default::default()
})),
}
}
}
struct OverlapResponder {
state: Mutex<OverlapState>,
both_entered: Condvar,
}
#[derive(Default)]
struct OverlapState {
active: usize,
max_active: usize,
arrivals: usize,
}
impl OverlapResponder {
fn new() -> Self {
Self {
state: Mutex::new(OverlapState::default()),
both_entered: Condvar::new(),
}
}
fn max_active(&self) -> usize {
self.state.lock().unwrap().max_active
}
}
impl HelloResponder for OverlapResponder {
fn handle_frame(&self, _frame: Frame, _peer: PeerIdentity) -> HelloReply {
let mut state = self.state.lock().unwrap();
state.active += 1;
state.arrivals += 1;
state.max_active = state.max_active.max(state.active);
self.both_entered.notify_all();
if state.arrivals < 2 {
state = self
.both_entered
.wait_timeout_while(state, std::time::Duration::from_millis(250), |state| {
state.arrivals < 2
})
.unwrap()
.0;
}
state.active -= 1;
HelloReply {
result: Some(HelloReplyResult::Negotiated(Negotiated::default())),
}
}
}
fn fixture_program() -> String {
let exe = std::env::current_exe().expect("test executable path");
let dir = exe
.parent()
.and_then(std::path::Path::parent)
.expect("test binary should live in <profile>/deps/");
dir.join(format!(
"testbin-stdio-scripted{}",
std::env::consts::EXE_SUFFIX
))
.to_string_lossy()
.into_owned()
}
fn frame(kind: session_frame::Kind) -> SessionFrame {
SessionFrame { kind: Some(kind) }
}
async fn read_framed_body<S: tokio::io::AsyncRead + Unpin>(stream: &mut S) -> Vec<u8> {
let mut version = [0u8; 1];
stream.read_exact(&mut version).await.expect("read version");
assert_eq!(version[0], ENVELOPE_VERSION, "framing version");
let mut len_buf = [0u8; 4];
stream.read_exact(&mut len_buf).await.expect("read len");
let len = u32::from_le_bytes(len_buf) as usize;
let mut body = vec![0u8; len];
stream.read_exact(&mut body).await.expect("read body");
body
}
#[cfg(unix)]
async fn send_hello_only(broker_path: &std::path::Path) {
let mut stream = IpcAsyncStream::connect(
&IpcEndpoint::new(broker_path.to_string_lossy().into_owned()).expect("client fs name"),
)
.await
.expect("client dials broker");
let hello_wire =
encode_framed(&Frame::request(CONTROL_PAYLOAD_PROTOCOL, Vec::new())).expect("encode hello");
stream.write_all(&hello_wire).await.expect("send hello");
let reply_body = read_framed_body(&mut stream).await;
let reply_frame = Frame::decode(reply_body.as_slice()).expect("decode reply frame");
let reply = HelloReply::decode(reply_frame.payload.as_slice()).expect("decode HelloReply");
assert!(matches!(
reply.result,
Some(HelloReplyResult::Negotiated(_))
));
}
#[cfg(unix)]
#[tokio::test]
async fn async_broker_negotiates_distinct_sessions_concurrently() {
let pid = std::process::id();
let broker_path = std::env::temp_dir().join(format!("rp-async-brk-overlap-{pid}.sock"));
let _ = std::fs::remove_file(&broker_path);
let broker_listener = AsyncListener::bind(
&IpcEndpoint::new(broker_path.as_path().to_string_lossy().into_owned())
.expect("broker fs name"),
)
.expect("bind broker endpoint");
let responder = Arc::new(OverlapResponder::new());
let broker_responder = Arc::clone(&responder);
let broker = tokio::spawn(async move {
let _ = serve_broker_session_endpoint_concurrently(
broker_listener,
broker_responder,
&PeerCredentialPolicy::allow_any(),
)
.await;
});
tokio::join!(send_hello_only(&broker_path), send_hello_only(&broker_path));
broker.abort();
let _ = std::fs::remove_file(&broker_path);
assert_eq!(
responder.max_active(),
2,
"one blocking route lookup must not serialize another SESSION Hello"
);
}
#[cfg(unix)]
#[tokio::test]
async fn async_broker_negotiates_hello_then_proxies_session() {
let pid = std::process::id();
let daemon_path = std::env::temp_dir().join(format!("rp-async-brk-d-{pid}.sock"));
let broker_path = std::env::temp_dir().join(format!("rp-async-brk-b-{pid}.sock"));
let _ = std::fs::remove_file(&daemon_path);
let _ = std::fs::remove_file(&broker_path);
let daemon_endpoint =
IpcEndpoint::new(daemon_path.to_string_lossy().into_owned()).expect("daemon endpoint");
let daemon_listener = AsyncListener::bind(&daemon_endpoint).expect("bind daemon endpoint");
let daemon = tokio::spawn(serve_session_endpoint(daemon_listener));
let broker_listener = AsyncListener::bind(
&IpcEndpoint::new(broker_path.as_path().to_string_lossy().into_owned())
.expect("broker fs name"),
)
.expect("bind broker endpoint");
let daemon_path_str = daemon_path.to_string_lossy().into_owned();
let broker = tokio::spawn(async move {
let _ = serve_broker_session_endpoint(
broker_listener,
&NegotiateToBackend(daemon_path_str),
&PeerCredentialPolicy::allow_any(),
)
.await;
});
let mut stream = IpcAsyncStream::connect(
&IpcEndpoint::new(broker_path.as_path().to_string_lossy().into_owned())
.expect("client fs name"),
)
.await
.expect("client dials broker");
let hello_wire =
encode_framed(&Frame::request(CONTROL_PAYLOAD_PROTOCOL, Vec::new())).expect("encode hello");
stream.write_all(&hello_wire).await.expect("send hello");
let reply_body = read_framed_body(&mut stream).await;
let reply_frame = Frame::decode(reply_body.as_slice()).expect("decode reply frame");
let reply = HelloReply::decode(reply_frame.payload.as_slice()).expect("decode HelloReply");
assert!(
matches!(reply.result, Some(HelloReplyResult::Negotiated(_))),
"broker negotiated the Hello"
);
let mut client = session_framed(stream);
client
.send(frame(session_frame::Kind::Start(SessionStart {
program: fixture_program(),
args: vec![
"out:HELLO".to_owned(),
"err:WORLD".to_owned(),
"exit:9".to_owned(),
],
cwd: String::new(),
env: Vec::new(),
clear_inherited_env: false,
environment_policy: 0,
})))
.await
.expect("send start");
client
.send(frame(session_frame::Kind::StdinEof(true)))
.await
.expect("send stdin eof");
let mut stdout = Vec::new();
let mut stderr = Vec::new();
let mut code = None;
while let Some(Ok(sf)) = client.next().await {
match sf.kind {
Some(session_frame::Kind::Stdout(b)) => stdout.extend_from_slice(&b),
Some(session_frame::Kind::Stderr(b)) => stderr.extend_from_slice(&b),
Some(session_frame::Kind::Exit(e)) => {
code = Some(e.code);
break;
}
_ => panic!("unexpected inbound-only frame on the outbound lane"),
}
}
daemon.abort();
broker.abort();
let _ = std::fs::remove_file(&daemon_path);
let _ = std::fs::remove_file(&broker_path);
assert_eq!(
stdout, b"HELLO",
"stdout proxied post-Hello across the broker"
);
assert_eq!(
stderr, b"WORLD",
"stderr proxied post-Hello across the broker"
);
assert_eq!(
code,
Some(9),
"exit code proxied after async Hello negotiation"
);
}
#[cfg(unix)]
#[tokio::test]
async fn async_broker_drops_peer_refused_by_policy() {
let pid = std::process::id();
let broker_path = std::env::temp_dir().join(format!("rp-async-brk-drop-{pid}.sock"));
let _ = std::fs::remove_file(&broker_path);
let broker_listener = AsyncListener::bind(
&IpcEndpoint::new(broker_path.as_path().to_string_lossy().into_owned())
.expect("broker fs name"),
)
.expect("bind broker endpoint");
let policy = PeerCredentialPolicy::owner_only("rp-no-such-owner-sentinel");
let broker = tokio::spawn(async move {
let _ = serve_broker_session_endpoint(
broker_listener,
&NegotiateToBackend(String::new()),
&policy,
)
.await;
});
let mut stream = IpcAsyncStream::connect(
&IpcEndpoint::new(broker_path.as_path().to_string_lossy().into_owned())
.expect("client fs name"),
)
.await
.expect("client dials broker");
let hello_wire =
encode_framed(&Frame::request(CONTROL_PAYLOAD_PROTOCOL, Vec::new())).expect("encode hello");
let _ = stream.write_all(&hello_wire).await;
let mut one = [0u8; 1];
let read = stream.read_exact(&mut one).await;
broker.abort();
let _ = std::fs::remove_file(&broker_path);
assert!(
read.is_err(),
"a peer refused by PeerCredentialPolicy must be dropped without a Hello reply"
);
}
#[cfg(unix)]
#[tokio::test]
async fn async_broker_session_socket_entry_binds_and_proxies() {
let pid = std::process::id();
let daemon_path = std::env::temp_dir().join(format!("rp-async-sock-d-{pid}.sock"));
let broker_path = std::env::temp_dir().join(format!("rp-async-sock-b-{pid}.sock"));
let _ = std::fs::remove_file(&daemon_path);
let _ = std::fs::remove_file(&broker_path);
let daemon_endpoint =
IpcEndpoint::new(daemon_path.to_string_lossy().into_owned()).expect("daemon endpoint");
let daemon_listener = AsyncListener::bind(&daemon_endpoint).expect("bind daemon endpoint");
let daemon = tokio::spawn(serve_session_endpoint(daemon_listener));
let daemon_path_str = daemon_path.to_string_lossy().into_owned();
let broker_path_str = broker_path.to_string_lossy().into_owned();
let broker = tokio::spawn(async move {
let _ = serve_broker_session_socket(
&broker_path_str,
&NegotiateToBackend(daemon_path_str),
&PeerCredentialPolicy::allow_any(),
)
.await;
});
let mut stream = None;
for _ in 0..200 {
match IpcAsyncStream::connect(
&IpcEndpoint::new(broker_path.as_path().to_string_lossy().into_owned())
.expect("client fs name"),
)
.await
{
Ok(s) => {
stream = Some(s);
break;
}
Err(_) => tokio::time::sleep(std::time::Duration::from_millis(10)).await,
}
}
let mut stream = stream.expect("broker socket became connectable");
let hello_wire =
encode_framed(&Frame::request(CONTROL_PAYLOAD_PROTOCOL, Vec::new())).expect("encode hello");
stream.write_all(&hello_wire).await.expect("send hello");
let reply_body = read_framed_body(&mut stream).await;
let reply_frame = Frame::decode(reply_body.as_slice()).expect("decode reply frame");
let reply = HelloReply::decode(reply_frame.payload.as_slice()).expect("decode HelloReply");
assert!(
matches!(reply.result, Some(HelloReplyResult::Negotiated(_))),
"broker negotiated the Hello via the socket entry"
);
let mut client = session_framed(stream);
client
.send(frame(session_frame::Kind::Start(SessionStart {
program: fixture_program(),
args: vec!["out:HELLO".to_owned(), "exit:7".to_owned()],
cwd: String::new(),
env: Vec::new(),
clear_inherited_env: false,
environment_policy: 0,
})))
.await
.expect("send start");
client
.send(frame(session_frame::Kind::StdinEof(true)))
.await
.expect("send stdin eof");
let mut stdout = Vec::new();
let mut code = None;
while let Some(Ok(sf)) = client.next().await {
match sf.kind {
Some(session_frame::Kind::Stdout(b)) => stdout.extend_from_slice(&b),
Some(session_frame::Kind::Stderr(_)) => {}
Some(session_frame::Kind::Exit(e)) => {
code = Some(e.code);
break;
}
_ => panic!("unexpected inbound-only frame on the outbound lane"),
}
}
daemon.abort();
broker.abort();
let _ = std::fs::remove_file(&daemon_path);
let _ = std::fs::remove_file(&broker_path);
assert_eq!(stdout, b"HELLO", "stdout proxied through the socket entry");
assert_eq!(code, Some(7), "exit code proxied through the socket entry");
}