use std::{
future::Future,
sync::{Arc, Mutex},
};
use http::uri::Authority;
use remoc::{self, RemoteSend, prelude::ServerShared};
use smallvec::smallvec;
use snafu::{ResultExt, Snafu};
use tokio_util::task::AbortOnDropHandle;
use tracing::Instrument;
use super::connection::{
ConnectionAdapter, ConnectionBootstrap, IPC_ERROR_KIND, IPC_FRAME_TYPE, IpcConnectionHandle,
IpcConnectionServerShared,
};
use crate::{
error::Code,
ipc::transport::{FdTransfer, MuxChannel, SplitError, TakeFdsError},
quic::{self, ConnectionError},
rpc::quic::serde_types::SerdeAuthority,
varint::VarInt,
};
#[remoc::rtc::remote]
pub trait IpcConnect: Send + Sync {
async fn connect(
&self,
server: SerdeAuthority,
fd_id: VarInt,
) -> Result<VarInt, ConnectionError>;
}
#[derive(Debug, Snafu)]
#[snafu(visibility(pub))]
pub enum ConnectorError {
#[snafu(display("ipc connect rpc failed"))]
Rpc { source: ConnectionError },
#[snafu(display("failed to retrieve mux channel fd"))]
WaitFds {
source: crate::ipc::transport::WaitFdsError,
},
#[snafu(display("unexpected connection fd count"))]
TakeFd { source: TakeFdsError },
#[snafu(display("ipc connect returned mismatched fd id {actual}, expected {expected}"))]
FdIdMismatch { expected: VarInt, actual: VarInt },
#[snafu(display("failed to reconstruct mux channel from fd"))]
FromFd { source: std::io::Error },
#[snafu(display("failed to split mux channel"))]
Split { source: SplitError },
#[snafu(display("failed to establish remoc connection: {message}"))]
Remoc { message: String },
#[snafu(display("failed to receive connection bootstrap"))]
Bootstrap,
}
pub struct ConnectAdapter<C, Codec> {
inner: C,
fd_transfer: FdTransfer,
tasks: Mutex<Vec<AbortOnDropHandle<()>>>,
_codec: std::marker::PhantomData<Codec>,
}
impl<C, Codec> ConnectAdapter<C, Codec> {
pub fn new(inner: C, fd_transfer: FdTransfer) -> Self {
Self {
inner,
fd_transfer,
tasks: Mutex::new(Vec::new()),
_codec: std::marker::PhantomData,
}
}
fn spawn_task(&self, task: impl Future<Output = ()> + Send + 'static) {
let handle = AbortOnDropHandle::new(tokio::spawn(task.in_current_span()));
let mut tasks = self
.tasks
.lock()
.expect("connect adapter task registry should not be poisoned");
tasks.retain(|task| !task.is_finished());
tasks.push(handle);
}
}
impl<C, Codec> IpcConnect for ConnectAdapter<C, Codec>
where
C: quic::Connect + 'static,
C::Connection: 'static,
<C::Connection as quic::ManageStream>::StreamReader: Unpin + 'static,
<C::Connection as quic::ManageStream>::StreamWriter: Unpin + 'static,
Codec: remoc::codec::Codec,
ConnectionBootstrap: RemoteSend,
{
async fn connect(
&self,
server: SerdeAuthority,
fd_id: VarInt,
) -> Result<VarInt, ConnectionError> {
let authority = Authority::try_from(server).map_err(|e| connect_error(e, "authority"))?;
let delivery = self
.fd_transfer
.delivery(fd_id)
.reserve()
.await
.map_err(|error| connect_error(error, "reserve fd delivery"))?;
let connection = quic::Connect::connect(&self.inner, &authority)
.await
.map_err(|e| connect_error(e, "connect"))?;
let (server_mux, client_fd) =
MuxChannel::create_pair().map_err(|e| connect_error(e, "create_pair"))?;
let (sink, stream) = match server_mux.split() {
Ok(split) => split,
Err(error) => {
close_undelivered_connection(connection.as_ref(), "split");
return Err(connect_error(error, "split"));
}
};
if let Err(error) = delivery.deliver(smallvec![client_fd]).await {
close_undelivered_connection(connection.as_ref(), "deliver");
return Err(connect_error(error, "deliver fd"));
}
self.spawn_task(Self::setup_connection(connection, sink, stream));
Ok(fd_id)
}
}
impl<C, Codec> ConnectAdapter<C, Codec>
where
C: quic::Connect + 'static,
C::Connection: 'static,
<C::Connection as quic::ManageStream>::StreamReader: Unpin + 'static,
<C::Connection as quic::ManageStream>::StreamWriter: Unpin + 'static,
Codec: remoc::codec::Codec,
ConnectionBootstrap: RemoteSend,
{
async fn setup_connection(
connection: Arc<C::Connection>,
sink: crate::ipc::transport::MuxSink,
stream: crate::ipc::transport::MuxStream,
) {
use tracing::debug;
let conn_fd_sender = sink.fd_sender();
let conn_fd_transfer = stream.fd_transfer(conn_fd_sender);
let (conn, mut tx, _rx) =
match remoc::Connect::framed::<_, _, ConnectionBootstrap, (), Codec>(
remoc::Cfg::default(),
sink,
stream,
)
.await
{
Ok(v) => v,
Err(e) => {
debug!(
error = %snafu::Report::from_error(e),
"per-connection remoc handshake failed"
);
return;
}
};
let remoc_task = AbortOnDropHandle::new(tokio::spawn(conn.in_current_span()));
let adapter = ConnectionAdapter::new(connection, conn_fd_transfer);
let (server, rpc_client) = IpcConnectionServerShared::new(Arc::new(adapter), 64);
let server_task = AbortOnDropHandle::new(tokio::spawn(
(async move {
let _ = server.serve(true).await;
})
.in_current_span(),
));
let bootstrap = ConnectionBootstrap {
connection: rpc_client,
};
if tx.send(bootstrap).await.is_err() {
debug!("failed to send connection bootstrap: base channel closed");
return;
}
let _ = futures::future::join(remoc_task, server_task).await;
}
}
pub struct IpcConnector<Codec> {
rpc: IpcConnectClient,
fd_transfer: FdTransfer,
_codec: std::marker::PhantomData<Codec>,
}
impl<Codec> IpcConnector<Codec> {
pub fn new(rpc: IpcConnectClient, fd_transfer: FdTransfer) -> Self {
Self {
rpc,
fd_transfer,
_codec: std::marker::PhantomData,
}
}
}
impl<Codec> quic::Connect for IpcConnector<Codec>
where
Codec: remoc::codec::Codec,
ConnectionBootstrap: RemoteSend,
{
type Connection = IpcConnectionHandle;
type Error = ConnectorError;
async fn connect<'a>(
&'a self,
server: &'a Authority,
) -> Result<Arc<IpcConnectionHandle>, ConnectorError> {
let receiver = self.fd_transfer.receive();
let fd_id = receiver.id();
let rpc = async {
let actual = IpcConnect::connect(&self.rpc, SerdeAuthority::from(server), fd_id)
.await
.context(RpcSnafu)?;
if actual != fd_id {
return Err(ConnectorError::FdIdMismatch {
expected: fd_id,
actual,
});
}
Ok(())
};
let receive = async { receiver.await.context(WaitFdsSnafu) };
let ((), received) = futures::future::try_join(rpc, receive).await?;
let fd = received.into_one().context(TakeFdSnafu)?;
let mux = MuxChannel::from_fd(fd).context(FromFdSnafu)?;
let (sink, stream) = mux.split().context(SplitSnafu)?;
let conn_fd_transfer = stream.fd_transfer(sink.fd_sender());
let (conn, _tx, mut rx) = remoc::Connect::framed::<_, _, (), ConnectionBootstrap, Codec>(
remoc::Cfg::default(),
sink,
stream,
)
.await
.map_err(|e| ConnectorError::Remoc {
message: e.to_string(),
})?;
let remoc_task = AbortOnDropHandle::new(tokio::spawn(
async move {
let _ = conn.await;
}
.in_current_span(),
));
let bootstrap = rx
.recv()
.await
.map_err(|_| ConnectorError::Bootstrap)?
.ok_or(ConnectorError::Bootstrap)?;
Ok(Arc::new(IpcConnectionHandle::new(
bootstrap.connection,
conn_fd_transfer,
remoc_task,
)))
}
}
fn connect_error(err: impl std::error::Error, context: &str) -> ConnectionError {
tracing::debug!(error = %snafu::Report::from_error(&err), context, "ipc connect error");
ConnectionError::Transport {
source: quic::TransportError {
kind: IPC_ERROR_KIND,
frame_type: IPC_FRAME_TYPE,
reason: format!("ipc connect: {context}").into(),
},
}
}
fn close_undelivered_connection(connection: &impl quic::Lifecycle, context: &'static str) {
quic::Lifecycle::close(
connection,
Code::H3_REQUEST_CANCELLED,
format!("ipc connect {context} failed").into(),
);
}