use std::future::Future;
use std::sync::Arc;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::sync::{mpsc, oneshot};
use super::protocol::{Answer, Compiled, Plan, Request, read_frame, write_frame};
use super::{ENDPOINT_ENV, Endpoint, TOKEN_ENV};
pub trait Handler: Send + Sync + 'static {
fn plan(
self: &Arc<Self>,
executable: std::ffi::OsString,
args: Vec<std::ffi::OsString>,
) -> impl Future<Output = Decision<Self::Pending>> + Send;
fn compiled(
self: &Arc<Self>,
pending: Self::Pending,
success: bool,
) -> impl Future<Output = ()> + Send;
type Pending: Send + 'static;
}
pub enum Decision<P> {
Served,
Compile(P),
}
pub struct Supervisor {
endpoint: Endpoint,
token: String,
shutdown: Option<oneshot::Sender<()>>,
#[cfg(unix)]
_socket_dir: Option<tempfile::TempDir>,
}
impl Supervisor {
#[must_use]
pub fn env(&self) -> [(&'static str, String); 2] {
[
(ENDPOINT_ENV, self.endpoint.encode()),
(TOKEN_ENV, self.token.clone()),
]
}
pub fn shutdown(&mut self) {
if let Some(shutdown) = self.shutdown.take() {
let _ = shutdown.send(());
}
}
}
impl Drop for Supervisor {
fn drop(&mut self) {
self.shutdown();
}
}
enum Ticket<P> {
Allocate(P, oneshot::Sender<u64>),
Take(u64, oneshot::Sender<Option<P>>),
}
pub fn start<H>(handler: Arc<H>) -> Result<Supervisor, String>
where
H: Handler,
{
let token = mint_token();
let (tickets, ticket_rx) = mpsc::unbounded_channel::<Ticket<H::Pending>>();
tokio::spawn(run_tickets(ticket_rx));
let (shutdown_tx, shutdown_rx) = oneshot::channel();
#[cfg(unix)]
{
let socket_dir = tempfile::Builder::new()
.prefix("stow-supervisor-")
.tempdir()
.map_err(|error| format!("create supervisor socket directory: {error}"))?;
let path = socket_dir.path().join("sock");
let listener = tokio::net::UnixListener::bind(&path)
.map_err(|error| format!("bind supervisor socket {}: {error}", path.display()))?;
restrict_socket(&path)?;
tokio::spawn(accept_unix(
listener,
handler,
tickets,
token.clone(),
shutdown_rx,
));
tracing::debug!(endpoint = %path.display(), "build supervisor listening");
Ok(Supervisor {
endpoint: Endpoint::Unix(path),
token,
shutdown: Some(shutdown_tx),
_socket_dir: Some(socket_dir),
})
}
#[cfg(not(unix))]
{
let listener = std::net::TcpListener::bind(("127.0.0.1", 0))
.map_err(|error| format!("bind supervisor loopback listener: {error}"))?;
listener
.set_nonblocking(true)
.map_err(|error| format!("set the supervisor listener non-blocking: {error}"))?;
let port = listener
.local_addr()
.map_err(|error| format!("read supervisor listener port: {error}"))?
.port();
let listener = tokio::net::TcpListener::from_std(listener)
.map_err(|error| format!("adopt the supervisor listener: {error}"))?;
tokio::spawn(accept_loopback(
listener,
handler,
tickets,
token.clone(),
shutdown_rx,
));
tracing::debug!(port, "build supervisor listening");
Ok(Supervisor {
endpoint: Endpoint::Loopback(port),
token,
shutdown: Some(shutdown_tx),
})
}
}
fn mint_token() -> String {
let mut bytes = [0u8; 16];
getrandom(&mut bytes);
hex::encode(bytes)
}
fn getrandom(bytes: &mut [u8; 16]) {
let seed = format!(
"{}-{}-{:?}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos(),
std::time::Instant::now()
);
bytes.copy_from_slice(&blake3::hash(seed.as_bytes()).as_bytes()[..16]);
}
#[cfg(unix)]
fn restrict_socket(path: &std::path::Path) -> Result<(), String> {
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))
.map_err(|error| format!("restrict supervisor socket {}: {error}", path.display()))
}
async fn run_tickets<P>(mut messages: mpsc::UnboundedReceiver<Ticket<P>>) {
let mut next = 1u64;
let mut pending = std::collections::HashMap::<u64, P>::new();
while let Some(message) = messages.recv().await {
match message {
Ticket::Allocate(value, reply) => {
let ticket = next;
next = next.wrapping_add(1);
pending.insert(ticket, value);
if reply.send(ticket).is_err() {
pending.remove(&ticket);
}
}
Ticket::Take(ticket, reply) => {
let _ = reply.send(pending.remove(&ticket));
}
}
}
}
#[cfg(unix)]
async fn accept_unix<H: Handler>(
listener: tokio::net::UnixListener,
handler: Arc<H>,
tickets: mpsc::UnboundedSender<Ticket<H::Pending>>,
token: String,
mut shutdown: oneshot::Receiver<()>,
) {
loop {
let accepted = tokio::select! {
() = async { (&mut shutdown).await.ok(); } => return,
accepted = listener.accept() => accepted,
};
match accepted {
Ok((stream, _)) => {
tokio::spawn(serve_connection(
stream,
Arc::clone(&handler),
tickets.clone(),
token.clone(),
));
}
Err(error) => {
tracing::warn!(%error, "supervisor accept failed");
return;
}
}
}
}
#[cfg(not(unix))]
async fn accept_loopback<H: Handler>(
listener: tokio::net::TcpListener,
handler: Arc<H>,
tickets: mpsc::UnboundedSender<Ticket<H::Pending>>,
token: String,
mut shutdown: oneshot::Receiver<()>,
) {
loop {
let accepted = tokio::select! {
() = async { (&mut shutdown).await.ok(); } => return,
accepted = listener.accept() => accepted,
};
match accepted {
Ok((stream, _)) => {
tokio::spawn(serve_connection(
stream,
Arc::clone(&handler),
tickets.clone(),
token.clone(),
));
}
Err(error) => {
tracing::warn!(%error, "supervisor accept failed");
return;
}
}
}
}
async fn serve_connection<S, H>(
mut stream: S,
handler: Arc<H>,
tickets: mpsc::UnboundedSender<Ticket<H::Pending>>,
token: String,
) where
S: AsyncRead + AsyncWrite + Unpin + Send,
H: Handler,
{
loop {
let request: Request = match read_frame(&mut stream).await {
Ok(Some(request)) => request,
Ok(None) => return,
Err(error) => {
tracing::warn!(%error, "supervisor could not read a facade frame");
return;
}
};
let answer = answer_request(&handler, &tickets, &token, request).await;
if let Err(error) = write_frame(&mut stream, &answer).await {
tracing::warn!(%error, "supervisor could not answer a facade");
return;
}
}
}
async fn answer_request<H: Handler>(
handler: &Arc<H>,
tickets: &mpsc::UnboundedSender<Ticket<H::Pending>>,
token: &str,
request: Request,
) -> Answer {
match request {
Request::Plan(plan) => answer_plan(handler, tickets, token, plan).await,
Request::Compiled(report) => answer_report(handler, tickets, token, report).await,
}
}
async fn answer_plan<H: Handler>(
handler: &Arc<H>,
tickets: &mpsc::UnboundedSender<Ticket<H::Pending>>,
token: &str,
plan: Plan,
) -> Answer {
if plan.token != token {
return Answer::Failed {
message: "supervisor token mismatch".to_owned(),
};
}
let (executable, args) = match (plan.executable(), plan.args()) {
(Ok(executable), Ok(args)) => (executable, args),
(Err(error), _) | (_, Err(error)) => return Answer::Failed { message: error },
};
match handler.plan(executable, args).await {
Decision::Served => Answer::Served,
Decision::Compile(pending) => {
let (sender, allocated) = oneshot::channel();
if tickets.send(Ticket::Allocate(pending, sender)).is_err() {
return Answer::Failed {
message: "supervisor ticket task is gone".to_owned(),
};
}
allocated.await.map_or_else(
|_| Answer::Failed {
message: "supervisor ticket task dropped the allocation".to_owned(),
},
|ticket| Answer::Compile { ticket },
)
}
}
}
async fn answer_report<H: Handler>(
handler: &Arc<H>,
tickets: &mpsc::UnboundedSender<Ticket<H::Pending>>,
token: &str,
report: Compiled,
) -> Answer {
if report.token != token {
return Answer::Failed {
message: "supervisor token mismatch".to_owned(),
};
}
let (sender, restored) = oneshot::channel();
if tickets.send(Ticket::Take(report.ticket, sender)).is_err() {
return Answer::Failed {
message: "supervisor ticket task is gone".to_owned(),
};
}
match restored.await {
Ok(Some(pending)) => {
handler.compiled(pending, report.success).await;
Answer::Recorded
}
Ok(None) => Answer::Failed {
message: format!("supervisor has no record of compile {}", report.ticket),
},
Err(_) => Answer::Failed {
message: "supervisor ticket task dropped the lookup".to_owned(),
},
}
}
#[cfg(test)]
mod tests {
use std::ffi::OsString;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use super::{Decision, Handler, start};
use crate::supervisor::client::{Connection, Decision as ClientDecision};
struct Stub {
reports: AtomicUsize,
}
impl Handler for Stub {
type Pending = OsString;
fn plan(
self: &Arc<Self>,
_executable: OsString,
args: Vec<OsString>,
) -> impl std::future::Future<Output = Decision<Self::Pending>> + Send {
std::future::ready(if args.first().is_some_and(|arg| arg == "--serve-me") {
Decision::Served
} else {
Decision::Compile(args.first().cloned().unwrap_or_default())
})
}
fn compiled(
self: &Arc<Self>,
_pending: Self::Pending,
success: bool,
) -> impl std::future::Future<Output = ()> + Send {
if success {
self.reports.fetch_add(1, Ordering::SeqCst);
}
std::future::ready(())
}
}
#[tokio::test]
async fn a_facade_plans_and_reports_over_the_wire() {
let handler = Arc::new(Stub {
reports: AtomicUsize::new(0),
});
let supervisor = start(Arc::clone(&handler)).expect("start the supervisor");
let [(_, endpoint), (_, token)] = supervisor.env();
let endpoint = crate::supervisor::Endpoint::parse(&endpoint).expect("endpoint");
let mut connection = Connection::open(&endpoint, token)
.await
.expect("connect to the supervisor");
let decision = connection
.plan(
std::ffi::OsStr::new("/usr/bin/rustc"),
&[OsString::from("--crate-name"), OsString::from("serde")],
)
.await
.expect("plan");
let ClientDecision::Compile(ticket) = decision else {
panic!("the stub asks for a compile");
};
connection.report(&ticket, true).await.expect("report");
assert_eq!(handler.reports.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn a_served_unit_ends_the_exchange() {
let handler = Arc::new(Stub {
reports: AtomicUsize::new(0),
});
let supervisor = start(Arc::clone(&handler)).expect("start the supervisor");
let [(_, endpoint), (_, token)] = supervisor.env();
let endpoint = crate::supervisor::Endpoint::parse(&endpoint).expect("endpoint");
let mut connection = Connection::open(&endpoint, token)
.await
.expect("connect to the supervisor");
let decision = connection
.plan(
std::ffi::OsStr::new("/usr/bin/rustc"),
&[OsString::from("--serve-me")],
)
.await
.expect("plan");
assert!(matches!(decision, ClientDecision::Served));
assert_eq!(handler.reports.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn a_wrong_token_is_refused() {
let handler = Arc::new(Stub {
reports: AtomicUsize::new(0),
});
let supervisor = start(handler).expect("start the supervisor");
let [(_, endpoint), _] = supervisor.env();
let endpoint = crate::supervisor::Endpoint::parse(&endpoint).expect("endpoint");
let mut connection = Connection::open(&endpoint, "not-the-token".to_owned())
.await
.expect("connect to the supervisor");
let error = connection
.plan(std::ffi::OsStr::new("/usr/bin/rustc"), &[])
.await
.expect_err("a wrong token must be refused");
assert!(error.contains("token mismatch"), "{error}");
}
}