use arcbox_docker::proxy::{GuestConnector, VsockShutdown, VsockStream};
use bytes::Bytes;
use hyper::body::Incoming;
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper::{Request, Response, StatusCode};
use hyper_util::rt::TokioIo;
use std::future::Future;
use std::path::{Path, PathBuf};
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use tokio::net::UnixListener;
use tokio::sync::Mutex;
use tokio_util::sync::CancellationToken;
#[derive(Clone)]
pub struct UnixSocketConnector {
pub socket_path: PathBuf,
pub connect_count: Arc<AtomicUsize>,
}
impl UnixSocketConnector {
pub fn new(socket_path: PathBuf) -> Self {
Self {
socket_path,
connect_count: Arc::new(AtomicUsize::new(0)),
}
}
}
impl GuestConnector for UnixSocketConnector {
fn connect(
&self,
) -> Pin<Box<dyn Future<Output = arcbox_docker::Result<TokioIo<VsockStream>>> + Send + '_>>
{
Box::pin(async {
self.connect_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let stream = tokio::net::UnixStream::connect(&self.socket_path)
.await
.map_err(|e| arcbox_docker::DockerError::Server(e.to_string()))?;
Ok(TokioIo::new(VsockStream::from_unix_stream_with_shutdown(
stream,
VsockShutdown::CloseOnDropOnly,
)))
})
}
}
pub struct MockRoute {
pub method: &'static str,
pub path: &'static str,
pub status: u16,
pub body: &'static str,
}
pub struct MockGuest {
pub socket_path: PathBuf,
#[allow(dead_code)]
pub cancel: CancellationToken,
last_upgrade_body: Arc<Mutex<Option<Bytes>>>,
last_request_uri: Arc<Mutex<Option<String>>>,
}
impl MockGuest {
#[allow(dead_code)] pub async fn last_upgrade_body(&self) -> Option<Bytes> {
self.last_upgrade_body.lock().await.clone()
}
#[allow(dead_code)] pub async fn last_request_uri(&self) -> Option<String> {
self.last_request_uri.lock().await.clone()
}
}
#[allow(dead_code)] pub async fn start(dir: &Path) -> MockGuest {
start_with_routes(dir, Vec::new()).await
}
pub async fn start_with_routes(dir: &Path, routes: Vec<MockRoute>) -> MockGuest {
let socket_path = dir.join("mock-guest.sock");
let listener = UnixListener::bind(&socket_path).expect("bind mock guest socket");
let cancel = CancellationToken::new();
let token = cancel.clone();
let last_upgrade_body: Arc<Mutex<Option<Bytes>>> = Arc::new(Mutex::new(None));
let body_slot = Arc::clone(&last_upgrade_body);
let last_request_uri: Arc<Mutex<Option<String>>> = Arc::new(Mutex::new(None));
let uri_slot = Arc::clone(&last_request_uri);
let routes = Arc::new(routes);
tokio::spawn(async move {
loop {
let stream = tokio::select! {
result = listener.accept() => match result {
Ok((s, _)) => s,
Err(_) => break,
},
() = token.cancelled() => break,
};
let slot = Arc::clone(&body_slot);
let uris = Arc::clone(&uri_slot);
let routes = Arc::clone(&routes);
tokio::spawn(async move {
let io = TokioIo::new(stream);
let svc = service_fn(move |req| {
let slot = Arc::clone(&slot);
let uris = Arc::clone(&uris);
let routes = Arc::clone(&routes);
handle(req, slot, uris, routes)
});
let _ = http1::Builder::new()
.serve_connection(io, svc)
.with_upgrades()
.await;
});
}
});
MockGuest {
socket_path,
cancel,
last_upgrade_body,
last_request_uri,
}
}
fn strip_version_prefix(path: &str) -> &str {
let Some(after_v) = path.strip_prefix("/v") else {
return path;
};
let Some(slash) = after_v.find('/') else {
return path;
};
let version = &after_v[..slash];
let (major, minor) = match version.split_once('.') {
Some(parts) => parts,
None => return path,
};
if !major.is_empty()
&& !minor.is_empty()
&& major.bytes().all(|b| b.is_ascii_digit())
&& minor.bytes().all(|b| b.is_ascii_digit())
{
&after_v[slash..]
} else {
path
}
}
async fn handle(
mut req: Request<Incoming>,
upgrade_body_slot: Arc<Mutex<Option<Bytes>>>,
request_uri_slot: Arc<Mutex<Option<String>>>,
routes: Arc<Vec<MockRoute>>,
) -> std::result::Result<Response<http_body_util::Full<Bytes>>, hyper::Error> {
*request_uri_slot.lock().await = Some(req.uri().to_string());
let path = strip_version_prefix(req.uri().path()).to_owned();
if let Some(route) = routes
.iter()
.find(|r| r.method == req.method().as_str() && r.path == path)
{
let resp = Response::builder()
.status(StatusCode::from_u16(route.status).expect("valid mock status"))
.header(hyper::header::CONTENT_TYPE, "application/json")
.body(http_body_util::Full::new(Bytes::from_static(
route.body.as_bytes(),
)))
.unwrap();
return Ok(resp);
}
let wants_upgrade = req.headers().contains_key(hyper::header::UPGRADE);
if wants_upgrade {
let upgrade_proto = req
.headers()
.get(hyper::header::UPGRADE)
.cloned()
.unwrap_or_else(|| hyper::header::HeaderValue::from_static("tcp"));
let body_bytes = http_body_util::BodyExt::collect(req.body_mut())
.await
.ok()
.map(|c| c.to_bytes())
.unwrap_or_default();
*upgrade_body_slot.lock().await = Some(body_bytes);
tokio::spawn(async move {
match hyper::upgrade::on(&mut req).await {
Ok(upgraded) => {
let mut io = TokioIo::new(upgraded);
let (mut rd, mut wr) = tokio::io::split(&mut io);
let _ = tokio::io::copy(&mut rd, &mut wr).await;
}
Err(e) => tracing::debug!("mock guest upgrade failed: {e}"),
}
});
let resp = Response::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.header(hyper::header::CONNECTION, "Upgrade")
.header(hyper::header::UPGRADE, upgrade_proto)
.body(http_body_util::Full::default())
.unwrap();
return Ok(resp);
}
let body_bytes = http_body_util::BodyExt::collect(req.into_body())
.await?
.to_bytes();
let resp = Response::builder()
.status(StatusCode::OK)
.header(hyper::header::CONTENT_TYPE, "application/json")
.body(http_body_util::Full::new(body_bytes))
.unwrap();
Ok(resp)
}