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}