use prost::Message;
use std::sync::Arc;
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use crate::broker::protocol::{
hello_reply::Result as HelloReplyResult, ErrorCode, Frame, FrameKind, HelloReply, Negotiated,
PayloadEncoding, CONTROL_PAYLOAD_PROTOCOL, ENVELOPE_VERSION, MAX_HELLO_BYTES, PROTOCOL_VERSION,
};
use crate::broker::session_relay::relay_local_socket_session;
use running_process_platform_internal::platform::ipc;
use super::connection::{
peer_identity_from_tokio_stream, refused_reply, HelloResponder, PeerCredentialPolicy,
};
use super::hello_handler::PeerIdentity;
const MAX_CONCURRENT_SESSION_NEGOTIATIONS: usize = 256;
pub async fn serve_broker_session_socket<R>(
socket_path: &str,
responder: &R,
peer_policy: &PeerCredentialPolicy,
) -> std::io::Result<()>
where
R: HelloResponder + ?Sized,
{
let listener = bind_session_listener(socket_path)?;
serve_broker_session_endpoint_opaque(listener, responder, peer_policy).await
}
fn bind_session_listener(socket_path: &str) -> std::io::Result<ipc::AsyncListener> {
let endpoint = ipc::Endpoint::new(socket_path.to_owned())?;
ipc::AsyncListener::bind(&endpoint)
}
pub async fn serve_broker_session_endpoint<R>(
listener: impl ipc::IntoAsyncListener,
responder: &R,
peer_policy: &PeerCredentialPolicy,
) -> std::io::Result<()>
where
R: HelloResponder + ?Sized,
{
serve_broker_session_endpoint_opaque(listener, responder, peer_policy).await
}
async fn serve_broker_session_endpoint_opaque<R>(
listener: impl ipc::IntoAsyncListener,
responder: &R,
peer_policy: &PeerCredentialPolicy,
) -> std::io::Result<()>
where
R: HelloResponder + ?Sized,
{
let listener = listener.into_async_listener();
loop {
let stream = listener.accept().await?;
let peer = match peer_identity_from_tokio_stream(&stream) {
Ok(peer) => peer,
Err(err) => {
eprintln!("running-process-broker: could not read peer credentials: {err}");
continue;
}
};
if !peer_policy.allows(&peer) {
eprintln!(
"running-process-broker: dropped session peer pid={} uid_or_sid={:?}: \
credential policy refused",
peer.pid, peer.uid_or_sid
);
continue;
}
match negotiate_session_hello(stream, responder, peer).await {
Ok(Some((stream, backend_pipe))) => {
if backend_pipe.is_empty() {
eprintln!(
"running-process-broker: negotiated a session but no backend endpoint \
is published; dropping"
);
continue;
}
tokio::spawn(async move {
if let Err(err) = relay_local_socket_session(stream, &backend_pipe).await {
eprintln!("running-process-broker: session relay ended: {err}");
}
});
}
Ok(None) => {}
Err(err) => {
eprintln!("running-process-broker: session hello failed: {err}");
}
}
}
}
pub async fn serve_broker_session_endpoint_concurrently<R>(
listener: impl ipc::IntoAsyncListener,
responder: Arc<R>,
peer_policy: &PeerCredentialPolicy,
) -> std::io::Result<()>
where
R: HelloResponder + Send + Sync + 'static,
{
serve_broker_session_endpoint_concurrently_opaque(listener, responder, peer_policy).await
}
async fn serve_broker_session_endpoint_concurrently_opaque<R>(
listener: impl ipc::IntoAsyncListener,
responder: Arc<R>,
peer_policy: &PeerCredentialPolicy,
) -> std::io::Result<()>
where
R: HelloResponder + Send + Sync + 'static,
{
let listener = listener.into_async_listener();
let permits = Arc::new(tokio::sync::Semaphore::new(
MAX_CONCURRENT_SESSION_NEGOTIATIONS,
));
loop {
let stream = listener.accept().await?;
let peer = match peer_identity_from_tokio_stream(&stream) {
Ok(peer) => peer,
Err(err) => {
eprintln!("running-process-broker: could not read peer credentials: {err}");
continue;
}
};
if !peer_policy.allows(&peer) {
eprintln!(
"running-process-broker: dropped session peer pid={} uid_or_sid={:?}: \
credential policy refused",
peer.pid, peer.uid_or_sid
);
continue;
}
let permit = Arc::clone(&permits)
.acquire_owned()
.await
.map_err(|_| std::io::Error::other("SESSION negotiation pool closed"))?;
let responder = Arc::clone(&responder);
tokio::spawn(async move {
let negotiated = {
let _permit = permit;
negotiate_session_hello_concurrently(stream, responder, peer).await
};
match negotiated {
Ok(Some((stream, backend_pipe))) => {
if backend_pipe.is_empty() {
eprintln!(
"running-process-broker: negotiated a session but no backend endpoint \
is published; dropping"
);
return;
}
if let Err(err) = relay_local_socket_session(stream, &backend_pipe).await {
eprintln!("running-process-broker: session relay ended: {err}");
}
}
Ok(None) => {}
Err(err) => {
eprintln!("running-process-broker: session hello failed: {err}");
}
}
});
}
}
async fn negotiate_session_hello_concurrently<S, R>(
mut stream: S,
responder: Arc<R>,
peer: PeerIdentity,
) -> std::io::Result<Option<(S, String)>>
where
S: AsyncRead + AsyncWrite + Unpin,
R: HelloResponder + Send + Sync + 'static,
{
let request_bytes = match read_hello_frame(&mut stream).await? {
Ok(bytes) => bytes,
Err(reply) => {
write_hello_response(&mut stream, None, &reply).await?;
return Ok(None);
}
};
let request_frame = match Frame::decode(request_bytes.as_slice()) {
Ok(frame) => frame,
Err(_) => {
let reply = refused_reply(ErrorCode::ErrorPeerRejected, "malformed broker Frame", 0);
write_hello_response(&mut stream, None, &reply).await?;
return Ok(None);
}
};
let route_frame = request_frame.clone();
let reply = tokio::task::spawn_blocking(move || responder.handle_frame(route_frame, peer))
.await
.map_err(|err| std::io::Error::other(format!("SESSION route worker failed: {err}")))?;
write_hello_response(&mut stream, Some(&request_frame), &reply).await?;
let backend_pipe = negotiated(&reply).map(|n| n.backend_pipe.clone());
Ok(backend_pipe.map(|pipe| (stream, pipe)))
}
pub async fn negotiate_session_hello<S, R>(
mut stream: S,
responder: &R,
peer: PeerIdentity,
) -> std::io::Result<Option<(S, String)>>
where
S: AsyncRead + AsyncWrite + Unpin,
R: HelloResponder + ?Sized,
{
let request_bytes = match read_hello_frame(&mut stream).await? {
Ok(bytes) => bytes,
Err(reply) => {
write_hello_response(&mut stream, None, &reply).await?;
return Ok(None);
}
};
let request_frame = match Frame::decode(request_bytes.as_slice()) {
Ok(frame) => frame,
Err(_) => {
let reply = refused_reply(ErrorCode::ErrorPeerRejected, "malformed broker Frame", 0);
write_hello_response(&mut stream, None, &reply).await?;
return Ok(None);
}
};
let reply = responder.handle_frame(request_frame.clone(), peer);
write_hello_response(&mut stream, Some(&request_frame), &reply).await?;
let backend_pipe = negotiated(&reply).map(|n| n.backend_pipe.clone());
Ok(backend_pipe.map(|pipe| (stream, pipe)))
}
fn negotiated(reply: &HelloReply) -> Option<&Negotiated> {
match reply.result.as_ref()? {
HelloReplyResult::Negotiated(n) => Some(n),
_ => None,
}
}
async fn read_hello_frame<S: AsyncRead + Unpin>(
stream: &mut S,
) -> std::io::Result<Result<Vec<u8>, HelloReply>> {
let mut version = [0u8; 1];
if read_exact_eof(stream, &mut version).await? {
return Ok(Err(refused_reply(
ErrorCode::ErrorPeerRejected,
"incomplete Hello frame",
0,
)));
}
if version[0] != ENVELOPE_VERSION {
return Ok(Err(refused_reply(
ErrorCode::ErrorVersionUnsupported,
"unsupported framing version",
0,
)));
}
let mut len_buf = [0u8; 4];
if read_exact_eof(stream, &mut len_buf).await? {
return Ok(Err(refused_reply(
ErrorCode::ErrorPeerRejected,
"incomplete Hello frame",
0,
)));
}
let body_len = u32::from_le_bytes(len_buf) as usize;
if body_len > MAX_HELLO_BYTES {
return Ok(Err(refused_reply(
ErrorCode::ErrorPeerRejected,
"initial Hello frame exceeds 64 KiB",
0,
)));
}
let mut body = vec![0u8; body_len];
if body_len > 0 && read_exact_eof(stream, &mut body).await? {
return Ok(Err(refused_reply(
ErrorCode::ErrorPeerRejected,
"incomplete Hello frame",
0,
)));
}
Ok(Ok(body))
}
async fn read_exact_eof<S: AsyncRead + Unpin>(
stream: &mut S,
buf: &mut [u8],
) -> std::io::Result<bool> {
match stream.read_exact(buf).await {
Ok(_) => Ok(false),
Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => Ok(true),
Err(err) => Err(err),
}
}
async fn write_hello_response<S: AsyncWrite + Unpin>(
stream: &mut S,
request_frame: Option<&Frame>,
reply: &HelloReply,
) -> std::io::Result<()> {
let response_frame = Frame {
envelope_version: PROTOCOL_VERSION,
kind: FrameKind::Response as i32,
payload_protocol: CONTROL_PAYLOAD_PROTOCOL,
payload: reply.encode_to_vec(),
request_id: request_frame.map_or(0, |frame| frame.request_id),
payload_encoding: PayloadEncoding::None as i32,
deadline_unix_ms: 0,
traceparent: request_frame
.map(|frame| frame.traceparent.clone())
.unwrap_or_default(),
tracestate: request_frame
.map(|frame| frame.tracestate.clone())
.unwrap_or_default(),
};
let body = response_frame.encode_to_vec();
let len = u32::try_from(body.len())
.map_err(|_| std::io::Error::other("broker response frame exceeds u32 length"))?;
let mut header = [0u8; 5];
header[0] = ENVELOPE_VERSION;
header[1..].copy_from_slice(&len.to_le_bytes());
stream.write_all(&header).await?;
if !body.is_empty() {
stream.write_all(&body).await?;
}
stream.flush().await?;
Ok(())
}
#[cfg(all(test, feature = "daemon"))]
mod tests;