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}