use std::sync::Arc;
use http::uri::Authority;
use remoc::{self, RemoteSend, prelude::ServerShared};
use smallvec::smallvec;
use snafu::{ResultExt, Snafu};
use tracing::Instrument;
use super::connection::{
ConnectionAdapter, ConnectionBootstrap, IPC_ERROR_KIND, IPC_FRAME_TYPE, IpcConnectionHandle,
IpcConnectionServerShared,
};
use crate::{
ipc::transport::{FdRegistry, FdSender, MuxChannel, SplitError},
quic::{self, ConnectionError},
rpc::quic::serde_types::SerdeAuthority,
varint::VarInt,
};
#[remoc::rtc::remote]
pub trait IpcConnect: Send + Sync {
async fn connect(&self, server: SerdeAuthority) -> 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("no fd received from registry"))]
EmptyFd,
#[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_sender: FdSender,
_codec: std::marker::PhantomData<Codec>,
}
impl<C, Codec> ConnectAdapter<C, Codec> {
pub fn new(inner: C, fd_sender: FdSender) -> Self {
Self {
inner,
fd_sender,
_codec: std::marker::PhantomData,
}
}
}
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) -> Result<VarInt, ConnectionError> {
let authority = Authority::try_from(server).map_err(|e| connect_error(e, "authority"))?;
let connection = quic::Connect::connect(&self.inner, &authority)
.await
.map_err(|e| connect_error(e, "connect"))?;
let connection = Arc::new(connection);
let (server_mux, client_fd) =
MuxChannel::create_pair().map_err(|e| connect_error(e, "create_pair"))?;
let fd_id = self
.fd_sender
.queue_fds(smallvec![client_fd])
.map_err(|e| connect_error(e, "queue_fds"))?;
let (sink, stream) = server_mux.split().map_err(|e| connect_error(e, "split"))?;
tokio::spawn(Self::setup_connection(connection, sink, stream).in_current_span());
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, 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;
}
};
tokio::spawn(conn.in_current_span());
let adapter = ConnectionAdapter::new(connection, conn_fd_sender);
let (server, rpc_client) = IpcConnectionServerShared::new(Arc::new(adapter), 64);
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");
}
}
}
pub struct IpcConnector<Codec> {
rpc: IpcConnectClient,
fd_registry: FdRegistry,
_codec: std::marker::PhantomData<Codec>,
}
impl<Codec> IpcConnector<Codec> {
pub fn new(rpc: IpcConnectClient, fd_registry: FdRegistry) -> Self {
Self {
rpc,
fd_registry,
_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<IpcConnectionHandle, ConnectorError> {
let fd_id = IpcConnect::connect(&self.rpc, SerdeAuthority::from(server))
.await
.context(RpcSnafu)?;
let fds = self
.fd_registry
.wait_fds(fd_id)
.await
.context(WaitFdsSnafu)?;
let fd = fds.into_iter().next().ok_or(ConnectorError::EmptyFd)?;
let mux = MuxChannel::from_fd(fd).context(FromFdSnafu)?;
let (sink, stream) = mux.split().context(SplitSnafu)?;
let conn_fd_sender = sink.fd_sender();
let conn_fd_registry = stream.fd_registry();
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(),
})?;
tokio::spawn(conn.in_current_span());
let bootstrap = rx
.recv()
.await
.map_err(|_| ConnectorError::Bootstrap)?
.ok_or(ConnectorError::Bootstrap)?;
Ok(IpcConnectionHandle::new(
bootstrap.connection,
conn_fd_registry,
conn_fd_sender,
))
}
}
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(),
},
}
}