git.delta.rocks / fleet / refs/commits / 14689aa0c750

difftreelog

source

cmds/usbd/src/server.rs5.7 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::{Utf8Path, 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: &Utf8Path, manifest: &Manifest) -> Result<Self> {43		let dir = dir.join(fleet_usb::DATA_DIR);44		let mut token = [0u8; 16];45		rand::rng().fill_bytes(&mut token);46		let token = hex::encode(token);4748		let store_dir = manifest49			.toplevel50			.parent()51			.context("toplevel should be a store path")?52			.to_string();53		let mut entries = HashMap::new();54		for entry in &manifest.paths {55			let base = entry56				.store_path57				.file_name()58				.context("store path should have a file name")?;59			let hash = base60				.split('-')61				.next()62				.expect("split returns at least one part");63			entries.insert(hash.to_owned(), entry.clone());64		}6566		let listener = TcpListener::bind(("127.0.0.1", 0))67			.await68			.context("binding local cache server")?;69		let port = listener.local_addr()?.port();70		let state = Arc::new(State {71			dir,72			prefix: format!("/{token}/"),73			store_dir,74			entries,75		});76		let accept_loop = tokio::spawn(accept_loop(listener, state));77		Ok(Self {78			url: format!("http://127.0.0.1:{port}/{token}"),79			accept_loop,80		})81	}82}8384async fn accept_loop(listener: TcpListener, state: Arc<State>) {85	loop {86		let stream = match listener.accept().await {87			Ok((stream, _)) => stream,88			Err(e) => {89				debug!("cache server accept: {e}");90				continue;91			}92		};93		let state = state.clone();94		tokio::spawn(async move {95			let service = service_fn(move |req| handle(state.clone(), req));96			if let Err(e) = hyper::server::conn::http1::Builder::new()97				.serve_connection(TokioIo::new(stream), service)98				.await99			{100				debug!("cache server connection: {e}");101			}102		});103	}104}105106async fn handle(state: Arc<State>, req: Request<Incoming>) -> Result<Response<Body>, Infallible> {107	let head = req.method() == Method::HEAD;108	if !head && req.method() != Method::GET {109		return Ok(status(StatusCode::METHOD_NOT_ALLOWED));110	}111	let mut resp = route(&state, req.uri().path()).unwrap_or_else(|e| {112		debug!("cache server request {} failed: {e}", req.uri().path());113		status(StatusCode::INTERNAL_SERVER_ERROR)114	});115	if head {116		*resp.body_mut() = empty();117	}118	Ok(resp)119}120121fn route(state: &Arc<State>, path: &str) -> Result<Response<Body>> {122	let Some(rest) = path.strip_prefix(&state.prefix) else {123		return Ok(status(StatusCode::NOT_FOUND));124	};125	if rest == "nix-cache-info" {126		return Ok(text(format!(127			"StoreDir: {}\nWantMassQuery: 1\nPriority: 30\n",128			state.store_dir129		)));130	}131	if let Some(hash) = rest.strip_suffix(".narinfo") {132		let Some(entry) = state.entries.get(hash) else {133			return Ok(status(StatusCode::NOT_FOUND));134		};135		return Ok(text(narinfo(hash, entry)));136	}137	if let Some(hash) = rest.strip_prefix("nar/") {138		let Some(entry) = state.entries.get(hash) else {139			return Ok(status(StatusCode::NOT_FOUND));140		};141		return nar(state, entry);142	}143	Ok(status(StatusCode::NOT_FOUND))144}145146fn narinfo(hash: &str, entry: &PathEntry) -> String {147	let references = entry148		.references149		.iter()150		.filter_map(|r| r.file_name())151		.collect::<Vec<_>>()152		.join(" ");153	let mut out = format!(154		"StorePath: {}\nURL: nar/{hash}\nCompression: {}\nNarHash: {}\nNarSize: {}\nReferences: {references}\n",155		entry.store_path, entry.compression, entry.nar_hash, entry.nar_size,156	);157	for sig in &entry.sigs {158		out.push_str(&format!("Sig: {sig}\n"));159	}160	out161}162163fn nar(state: &Arc<State>, entry: &PathEntry) -> Result<Response<Body>> {164	let key: [u8; 32] = BASE64165		.decode(&entry.key)?166		.try_into()167		.map_err(|_| anyhow!("invalid key size in manifest"))?;168	let files = entry169		.chunks170		.iter()171		.map(|c| state.dir.join(fleet_usb::names::data_rel_path(c)))172		.collect::<Vec<_>>();173174	let (mut tx, rx) = futures::channel::mpsc::channel::<Result<Frame<Bytes>, std::io::Error>>(8);175	spawn_blocking(move || {176		let mut stream = || -> std::io::Result<()> {177			let mut chunks: Box<dyn std::io::Read> = Box::new(std::io::empty());178			for file in &files {179				chunks = Box::new(chunks.chain(File::open(file)?));180			}181			let mut reader = DecryptReader::new(&key, chunks);182			let mut buf = vec![0u8; 64 * 1024];183			loop {184				let n = reader.read(&mut buf)?;185				if n == 0 {186					return Ok(());187				}188				let frame = Frame::data(Bytes::copy_from_slice(&buf[..n]));189				if futures::executor::block_on(tx.send(Ok(frame))).is_err() {190					return Ok(());191				}192			}193		};194		if let Err(e) = stream() {195			let _ = futures::executor::block_on(tx.send(Err(e)));196		}197	});198	Ok(Response::new(StreamBody::new(rx).boxed()))199}200201fn empty() -> Body {202	Full::new(Bytes::new())203		.map_err(|never| match never {})204		.boxed()205}206207fn text(s: String) -> Response<Body> {208	Response::new(209		Full::new(Bytes::from(s))210			.map_err(|never| match never {})211			.boxed(),212	)213}214215fn status(code: StatusCode) -> Response<Body> {216	let mut resp = Response::new(empty());217	*resp.status_mut() = code;218	resp219}