git.delta.rocks / fleet / refs/commits / 615754ca0747

difftreelog

source

cmds/usbd/src/server.rs5.6 KiBsourcehistory
1use std::{collections::HashMap, convert::Infallible, fs::File, io::Read as _, sync::Arc};23use anyhow::{Context as _, Result, anyhow};4use base64::Engine as _;5use base64::engine::general_purpose::STANDARD as BASE64;6use bytes::Bytes;7use camino::Utf8PathBuf;8use fleet_usb::{manifest::Manifest, manifest::PathEntry, stream::DecryptReader};9use futures::SinkExt as _;10use http_body_util::{BodyExt as _, Full, StreamBody, combinators::BoxBody};11use hyper::{12	Method, Request, Response, StatusCode,13	body::{Frame, Incoming},14	service::service_fn,15};16use hyper_util::rt::TokioIo;17use rand::Rng as _;18use tokio::{net::TcpListener, task::spawn_blocking};19use tracing::debug;2021type Body = BoxBody<Bytes, std::io::Error>;2223struct State {24	dir: Utf8PathBuf,25	prefix: String,26	store_dir: String,27	entries: HashMap<String, PathEntry>,28}2930pub struct CacheServer {31	pub url: String,32	accept_loop: tokio::task::JoinHandle<()>,33}3435impl Drop for CacheServer {36	fn drop(&mut self) {37		self.accept_loop.abort();38	}39}4041impl CacheServer {42	pub async fn start(dir: Utf8PathBuf, manifest: &Manifest) -> Result<Self> {43		let mut token = [0u8; 16];44		rand::rng().fill_bytes(&mut token);45		let token = hex::encode(token);4647		let store_dir = manifest48			.toplevel49			.parent()50			.context("toplevel should be a store path")?51			.to_string();52		let mut entries = HashMap::new();53		for entry in &manifest.paths {54			let base = entry55				.store_path56				.file_name()57				.context("store path should have a file name")?;58			let hash = base59				.split('-')60				.next()61				.expect("split returns at least one part");62			entries.insert(hash.to_owned(), entry.clone());63		}6465		let listener = TcpListener::bind(("127.0.0.1", 0))66			.await67			.context("binding local cache server")?;68		let port = listener.local_addr()?.port();69		let state = Arc::new(State {70			dir,71			prefix: format!("/{token}/"),72			store_dir,73			entries,74		});75		let accept_loop = tokio::spawn(accept_loop(listener, state));76		Ok(Self {77			url: format!("http://127.0.0.1:{port}/{token}"),78			accept_loop,79		})80	}81}8283async fn accept_loop(listener: TcpListener, state: Arc<State>) {84	loop {85		let stream = match listener.accept().await {86			Ok((stream, _)) => stream,87			Err(e) => {88				debug!("cache server accept: {e}");89				continue;90			}91		};92		let state = state.clone();93		tokio::spawn(async move {94			let service = service_fn(move |req| handle(state.clone(), req));95			if let Err(e) = hyper::server::conn::http1::Builder::new()96				.serve_connection(TokioIo::new(stream), service)97				.await98			{99				debug!("cache server connection: {e}");100			}101		});102	}103}104105async fn handle(state: Arc<State>, req: Request<Incoming>) -> Result<Response<Body>, Infallible> {106	let head = req.method() == Method::HEAD;107	if !head && req.method() != Method::GET {108		return Ok(status(StatusCode::METHOD_NOT_ALLOWED));109	}110	let mut resp = route(&state, req.uri().path()).unwrap_or_else(|e| {111		debug!("cache server request {} failed: {e}", req.uri().path());112		status(StatusCode::INTERNAL_SERVER_ERROR)113	});114	if head {115		*resp.body_mut() = empty();116	}117	Ok(resp)118}119120fn route(state: &Arc<State>, path: &str) -> Result<Response<Body>> {121	let Some(rest) = path.strip_prefix(&state.prefix) else {122		return Ok(status(StatusCode::NOT_FOUND));123	};124	if rest == "nix-cache-info" {125		return Ok(text(format!(126			"StoreDir: {}\nWantMassQuery: 1\nPriority: 30\n",127			state.store_dir128		)));129	}130	if let Some(hash) = rest.strip_suffix(".narinfo") {131		let Some(entry) = state.entries.get(hash) else {132			return Ok(status(StatusCode::NOT_FOUND));133		};134		return Ok(text(narinfo(hash, entry)));135	}136	if let Some(hash) = rest.strip_prefix("nar/") {137		let Some(entry) = state.entries.get(hash) else {138			return Ok(status(StatusCode::NOT_FOUND));139		};140		return nar(state, entry);141	}142	Ok(status(StatusCode::NOT_FOUND))143}144145fn narinfo(hash: &str, entry: &PathEntry) -> String {146	let references = entry147		.references148		.iter()149		.filter_map(|r| r.file_name())150		.collect::<Vec<_>>()151		.join(" ");152	let mut out = format!(153		"StorePath: {}\nURL: nar/{hash}\nCompression: {}\nNarHash: {}\nNarSize: {}\nReferences: {references}\n",154		entry.store_path, entry.compression, entry.nar_hash, entry.nar_size,155	);156	for sig in &entry.sigs {157		out.push_str(&format!("Sig: {sig}\n"));158	}159	out160}161162fn nar(state: &Arc<State>, entry: &PathEntry) -> Result<Response<Body>> {163	let key: [u8; 32] = BASE64164		.decode(&entry.key)?165		.try_into()166		.map_err(|_| anyhow!("invalid key size in manifest"))?;167	let files = entry168		.chunks169		.iter()170		.map(|c| state.dir.join(c))171		.collect::<Vec<_>>();172173	let (mut tx, rx) = futures::channel::mpsc::channel::<Result<Frame<Bytes>, std::io::Error>>(8);174	spawn_blocking(move || {175		let mut stream = || -> std::io::Result<()> {176			let mut chunks: Box<dyn std::io::Read> = Box::new(std::io::empty());177			for file in &files {178				chunks = Box::new(chunks.chain(File::open(file)?));179			}180			let mut reader = DecryptReader::new(&key, chunks);181			let mut buf = vec![0u8; 64 * 1024];182			loop {183				let n = reader.read(&mut buf)?;184				if n == 0 {185					return Ok(());186				}187				let frame = Frame::data(Bytes::copy_from_slice(&buf[..n]));188				if futures::executor::block_on(tx.send(Ok(frame))).is_err() {189					return Ok(());190				}191			}192		};193		if let Err(e) = stream() {194			let _ = futures::executor::block_on(tx.send(Err(e)));195		}196	});197	Ok(Response::new(StreamBody::new(rx).boxed()))198}199200fn empty() -> Body {201	Full::new(Bytes::new())202		.map_err(|never| match never {})203		.boxed()204}205206fn text(s: String) -> Response<Body> {207	Response::new(208		Full::new(Bytes::from(s))209			.map_err(|never| match never {})210			.boxed(),211	)212}213214fn status(code: StatusCode) -> Response<Body> {215	let mut resp = Response::new(empty());216	*resp.status_mut() = code;217	resp218}