diff --git a/CHANGELOG.md b/CHANGELOG.md index 17a3d0c85..cfd2eb3ca 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,7 +2,7 @@ - **Fixed** An invalid glob in `--filter` no longer shows its error message twice ([#763](https://github.com/voidzero-dev/vite-task/pull/763)). - **Changed** The detailed summary from `vp run --verbose` and `vp run --last-details` now shows each underlying cause of an error on its own line ([#761](https://github.com/voidzero-dev/vite-task/pull/761)). -- **Added** Remote caching. Configure an endpoint with the workspace's `cache: { remote: { url } }` or `VP_REMOTE_CACHE_URL`, and choose access with `--remote-cache=off|read|read-write` or `VP_REMOTE_CACHE`. The default is `read` with an endpoint and `off` without one. After a local cache miss, `vp run` looks the task up in the remote cache and, on a hit, restores its outputs and caches it locally. The task output and the run summary show which hits came from the remote cache. A failed read is just a cache miss, with the failure as its reason. In `read-write` mode, `vp run` also uploads the results of successful, cacheable tasks after caching them locally. A failed upload doesn't fail the task; the run summary shows a warning instead. Ctrl-C, or a failing task, stops remote cache requests right away, and a task still being looked up doesn't start. Tasks can opt out with `cache: { remote: false }`. Requests use the proxy environment variables or, on macOS and Windows, the system proxy settings ([#727](https://github.com/voidzero-dev/vite-task/pull/727), [#755](https://github.com/voidzero-dev/vite-task/pull/755), [#756](https://github.com/voidzero-dev/vite-task/pull/756), [#757](https://github.com/voidzero-dev/vite-task/pull/757), [#764](https://github.com/voidzero-dev/vite-task/pull/764), [#771](https://github.com/voidzero-dev/vite-task/pull/771), [#772](https://github.com/voidzero-dev/vite-task/pull/772)). +- **Added** Remote caching. Configure an endpoint with the workspace's `cache: { remote: { url } }` or `VP_REMOTE_CACHE_URL`, and choose access with `--remote-cache=off|read|read-write` or `VP_REMOTE_CACHE`. The default is `read` with an endpoint and `off` without one. After a local cache miss, `vp run` looks the task up in the remote cache and, on a hit, restores its outputs and caches it locally. The task output and the run summary show which hits came from the remote cache. A failed read is just a cache miss, with the failure as its reason. In `read-write` mode, `vp run` also uploads the results of successful, cacheable tasks after caching them locally. A failed upload doesn't fail the task; the run summary shows a warning instead. In a GitHub Actions job granted `id-token: write`, uploads authenticate with a GitHub OIDC token that `vp run` requests for the endpoint, so no token needs to be configured. Once an upload is rejected as unauthorized, or no token can be obtained, `vp run` stops uploading to that endpoint for the rest of the run. Ctrl-C, or a failing task, stops remote cache requests right away, and a task still being looked up doesn't start. Tasks can opt out with `cache: { remote: false }`. Requests use the proxy environment variables or, on macOS and Windows, the system proxy settings ([#727](https://github.com/voidzero-dev/vite-task/pull/727), [#755](https://github.com/voidzero-dev/vite-task/pull/755), [#756](https://github.com/voidzero-dev/vite-task/pull/756), [#757](https://github.com/voidzero-dev/vite-task/pull/757), [#764](https://github.com/voidzero-dev/vite-task/pull/764), [#771](https://github.com/voidzero-dev/vite-task/pull/771), [#772](https://github.com/voidzero-dev/vite-task/pull/772), [#774](https://github.com/voidzero-dev/vite-task/pull/774)). - **Fixed** On Windows, environment variable names used by `vp run` now match regardless of ASCII letter case. Assignments in task commands override earlier assignments and inherited variables spelled differently, and `FORCE_COLOR`, `VP_RUN_CONCURRENCY_LIMIT`, and variables requested through `@voidzero-dev/vite-task-client` are found under any spelling ([#747](https://github.com/voidzero-dev/vite-task/pull/747)). - **Changed** A task's cache settings now go inside `cache`, e.g. `cache: { env: ["NODE_ENV"], input: ["src/**"] }`; `cache: true` is the same as `cache: {}`. `env`, `untrackedEnv`, `input`, and `output` are no longer supported at the top level of a task ([#749](https://github.com/voidzero-dev/vite-task/pull/749)). - **Fixed** Cached tasks on macOS no longer intermittently fail with exit 2 and `oils I/O error (main): No such process` when a fast command finishes before the shell gets scheduled. The bundled shell that runs task commands is updated to Oils 0.38.0, which fixes this race ([#702](https://github.com/voidzero-dev/vite-task/issues/702), [#703](https://github.com/voidzero-dev/vite-task/pull/703)). diff --git a/Cargo.lock b/Cargo.lock index 67a03f1e3..625741d0b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5250,12 +5250,14 @@ dependencies = [ name = "vt_remote_cache" version = "0.0.0" dependencies = [ + "base64 0.22.1", "bytes", "ciborium", "reqwest", "rustls", "serde", "serde_bytes", + "serde_json", "tempfile", "thiserror 2.0.18", "tokio", diff --git a/crates/vt/src/session/cache/mod.rs b/crates/vt/src/session/cache/mod.rs index 1f2d725ef..1a1a4ce9b 100644 --- a/crates/vt/src/session/cache/mod.rs +++ b/crates/vt/src/session/cache/mod.rs @@ -23,6 +23,7 @@ use vt_plan::{ cache_metadata::{CacheMetadata, ExecutionCacheKey, SpawnFingerprint}, remote_cache::{RemoteCacheAccess, ResolvedRemoteCacheConfig}, }; +use vt_remote_cache::StoreAuth; use vt_str::Str; use wincode::{ SchemaRead, SchemaReadOwned, SchemaWrite, @@ -312,8 +313,10 @@ pub fn cache_schema_dir_name() -> Str { } impl ExecutionCache { + /// Open the cache in `path`. Uploads to the remote cache authenticate with + /// `store_auth`. #[tracing::instrument(level = "debug", skip_all)] - pub fn load_from_path(path: &AbsolutePath) -> anyhow::Result { + pub fn load_from_path(path: &AbsolutePath, store_auth: StoreAuth) -> anyhow::Result { tracing::info!("Creating task cache directory at {}", path.as_path().display()); std::fs::create_dir_all(path)?; @@ -337,7 +340,7 @@ impl ExecutionCache { CREATE TABLE IF NOT EXISTS task_fingerprints (key BLOB PRIMARY KEY, value BLOB);", )?; // Lock is released when lock_file is dropped - Ok(Self { conn: Mutex::new(conn), remote_clients: RemoteClients::default() }) + Ok(Self { conn: Mutex::new(conn), remote_clients: RemoteClients::new(store_auth) }) } #[tracing::instrument] @@ -711,7 +714,7 @@ mod tests { fn reopening_preserves_existing_entries() { let (_tmp, dir) = temp_dir(); - drop(ExecutionCache::load_from_path(&dir).unwrap()); + drop(ExecutionCache::load_from_path(&dir, StoreAuth::Anonymous).unwrap()); { let conn = open_raw(&dir.join("cache.db")); conn.execute("INSERT INTO cache_entries (key, value) VALUES (X'01', X'02')", ()) @@ -719,7 +722,7 @@ mod tests { } // Reopening must not recreate or clear the tables. - drop(ExecutionCache::load_from_path(&dir).unwrap()); + drop(ExecutionCache::load_from_path(&dir, StoreAuth::Anonymous).unwrap()); let count: u32 = open_raw(&dir.join("cache.db")) .query_one("SELECT COUNT(*) FROM cache_entries", (), |r| r.get(0)) @@ -738,8 +741,8 @@ mod tests { let dir_a = base.join("v13"); let dir_b = base.join("v14"); - drop(ExecutionCache::load_from_path(&dir_a).unwrap()); - drop(ExecutionCache::load_from_path(&dir_b).unwrap()); + drop(ExecutionCache::load_from_path(&dir_a, StoreAuth::Anonymous).unwrap()); + drop(ExecutionCache::load_from_path(&dir_b, StoreAuth::Anonymous).unwrap()); assert!(dir_a.join("cache.db").as_path().exists()); assert!(dir_b.join("cache.db").as_path().exists()); diff --git a/crates/vt/src/session/cache/remote.rs b/crates/vt/src/session/cache/remote.rs index 18008f28e..6ce84b2b7 100644 --- a/crates/vt/src/session/cache/remote.rs +++ b/crates/vt/src/session/cache/remote.rs @@ -14,18 +14,20 @@ //! their keys differ. use std::{ + ffi::OsStr, fs::File, io::{self, Write as _}, - sync::{Arc, Mutex, PoisonError}, + sync::{Arc, Mutex, OnceLock, PoisonError}, }; use bytes::Bytes; use rustc_hash::FxHashMap; use tokio::sync::mpsc; use tokio_util::sync::CancellationToken; +use vt_casefold::EnvName; use vt_path::AbsolutePath; use vt_plan::cache_metadata::ExecutionCacheKey; -use vt_remote_cache::{Client, Download, Fetched}; +use vt_remote_cache::{Client, Download, Fetched, GithubOidc, StoreAuth}; use vt_str::Str; use wincode::{ SchemaWrite, @@ -42,6 +44,12 @@ use super::{ pub enum UploadError { #[error(transparent)] Remote(#[from] vt_remote_cache::Error), + /// Uploads to the endpoint aren't authorized: no GitHub Actions OIDC token + /// could be obtained, or the server responded with 401 or 403. Uploads to + /// the endpoint stop, and each later one returns the error that stopped + /// them without a request, so they all report the same reason. + #[error(transparent)] + Unauthorized(Arc), #[error("failed to encode the cache entry")] Encode(#[from] WriteError), /// The run was cancelled, by Ctrl-C or fast-fail, before the upload @@ -131,22 +139,58 @@ pub(super) fn resolve( } } +/// How uploads authenticate, from the session envs. A GitHub Actions job +/// granted `id-token: write` has both OIDC token request variables, and its +/// uploads send a token. These variables are also passed through to tasks, +/// so tools such as npm can use them too. +pub fn store_auth(envs: &FxHashMap>, Arc>) -> StoreAuth { + let env = |name: &str| { + envs.get(EnvName::from_ref(OsStr::new(name))) + .and_then(|value| value.to_str()) + .filter(|value| !value.is_empty()) + }; + match (env("ACTIONS_ID_TOKEN_REQUEST_URL"), env("ACTIONS_ID_TOKEN_REQUEST_TOKEN")) { + (Some(request_url), Some(request_token)) => { + StoreAuth::GithubOidc(GithubOidc::new(request_url, request_token)) + } + _ if env("GITHUB_ACTIONS") == Some("true") => StoreAuth::GithubActionsWithoutOidc, + _ => StoreAuth::Anonymous, + } +} + /// Remote cache clients, each created when its endpoint is first used. #[derive(Debug, Default)] pub struct RemoteClients { - clients: Mutex, Arc>>, + store_auth: StoreAuth, + endpoints: Mutex, Arc>>, +} + +#[derive(Debug)] +struct Endpoint { + client: Client, + /// The error that stopped uploads to the endpoint. See + /// [`UploadError::Unauthorized`]. + uploads_stopped: OnceLock>, } impl RemoteClients { - fn client(&self, endpoint: &Arc) -> Result, vt_remote_cache::Error> { - let mut clients = self.clients.lock().unwrap_or_else(PoisonError::into_inner); - if let Some(client) = clients.get(endpoint) { - return Ok(Arc::clone(client)); + /// Clients whose uploads authenticate with `store_auth`. + pub fn new(store_auth: StoreAuth) -> Self { + Self { store_auth, endpoints: Mutex::default() } + } + + fn endpoint(&self, endpoint: &Arc) -> Result, vt_remote_cache::Error> { + let mut endpoints = self.endpoints.lock().unwrap_or_else(PoisonError::into_inner); + if let Some(existing) = endpoints.get(endpoint) { + return Ok(Arc::clone(existing)); } - let client = Arc::new(Client::new(endpoint)?); - clients.insert(Arc::clone(endpoint), Arc::clone(&client)); - drop(clients); - Ok(client) + let created = Arc::new(Endpoint { + client: Client::new(endpoint, self.store_auth.clone())?, + uploads_stopped: OnceLock::new(), + }); + endpoints.insert(Arc::clone(endpoint), Arc::clone(&created)); + drop(endpoints); + Ok(created) } /// Fetch the entry stored under `cache_key`, falling back to the entry @@ -159,11 +203,11 @@ impl RemoteClients { execution_cache_key: &ExecutionCacheKey, cancel_token: &CancellationToken, ) -> Result, ReadError> { - let client = self.client(endpoint).map_err(ReadError::Fetch)?; + let endpoint = self.endpoint(endpoint).map_err(ReadError::Fetch)?; let key = encode_key(cache_key)?; let secondary_key = encode_key(execution_cache_key)?; cancel_token - .run_until_cancelled(client.fetch(&key, &secondary_key)) + .run_until_cancelled(endpoint.client.fetch(&key, &secondary_key)) .await .ok_or(ReadError::Cancelled)? .map_err(ReadError::Fetch) @@ -180,12 +224,13 @@ impl RemoteClients { cache_dir: &AbsolutePath, cancel_token: &CancellationToken, ) -> Result { - let client = self.client(endpoint).map_err(ReadError::Download)?; + let endpoint = self.endpoint(endpoint).map_err(ReadError::Download)?; let archive_name = vt_str::format!("{}.tar.zst", uuid::Uuid::new_v4()); let archive_path = cache_dir.join(archive_name.as_str()); let temp_path = cache_dir.join(vt_str::format!("{archive_name}.tmp").as_str()); - let result = - download_checked(&client, blob_id, &temp_path, cancel_token).await.and_then(|()| { + let result = download_checked(&endpoint.client, blob_id, &temp_path, cancel_token) + .await + .and_then(|()| { std::fs::rename(temp_path.as_path(), archive_path.as_path()) .map_err(ReadError::WriteArchive) }); @@ -197,7 +242,9 @@ impl RemoteClients { } /// Upload an entry that was just recorded locally, along with its output - /// archive in `cache_dir`. Stops when `cancel_token` is cancelled. + /// archive in `cache_dir`. Stops when `cancel_token` is cancelled. Once an + /// upload to the endpoint is unauthorized, later ones return its error + /// without a request. pub(super) async fn upload( &self, endpoint: &Arc, @@ -207,14 +254,26 @@ impl RemoteClients { cache_dir: &AbsolutePath, cancel_token: &CancellationToken, ) -> Result<(), UploadError> { - let client = self.client(endpoint)?; + let endpoint = self.endpoint(endpoint)?; + if let Some(err) = endpoint.uploads_stopped.get() { + return Err(UploadError::Unauthorized(Arc::clone(err))); + } let key = encode_key(cache_key)?; let secondary_key = encode_key(execution_cache_key)?; let value = serialize_cache(cache_value)?; let archive = cache_value.output_archive.as_ref().map(|name| cache_dir.join(name.as_str())); - let store = client.store(&key, &secondary_key, &value, archive.as_deref()); - cancel_token.run_until_cancelled(store).await.ok_or(UploadError::Cancelled)??; - Ok(()) + let store = endpoint.client.store(&key, &secondary_key, &value, archive.as_deref()); + match cancel_token.run_until_cancelled(store).await.ok_or(UploadError::Cancelled)? { + Ok(()) => Ok(()), + Err(err) if err.is_unauthorized() => { + let err = Arc::new(err); + // Concurrent uploads may fail the same way. The first to + // finish stops the rest. + let _ = endpoint.uploads_stopped.set(Arc::clone(&err)); + Err(UploadError::Unauthorized(err)) + } + Err(err) => Err(err.into()), + } } } @@ -633,6 +692,174 @@ mod tests { assert_eq!(std::fs::read_dir(dir.path()).unwrap().count(), 0); } + fn envs(pairs: &[(&str, &str)]) -> FxHashMap>, Arc> { + pairs + .iter() + .map(|(name, value)| { + (EnvName::new(Arc::::from(OsStr::new(name))), Arc::from(OsStr::new(value))) + }) + .collect() + } + + #[test] + fn store_auth_uses_oidc_when_both_variables_are_set() { + let url = ("ACTIONS_ID_TOKEN_REQUEST_URL", "https://token.example/?api-version=2.0"); + let token = ("ACTIONS_ID_TOKEN_REQUEST_TOKEN", "request-token"); + let github_actions = ("GITHUB_ACTIONS", "true"); + for (pairs, expected) in [ + (vec![url, token], "GithubOidc"), + (vec![url, token, github_actions], "GithubOidc"), + ( + vec![url, ("ACTIONS_ID_TOKEN_REQUEST_TOKEN", ""), github_actions], + "GithubActionsWithoutOidc", + ), + (vec![token, github_actions], "GithubActionsWithoutOidc"), + (vec![github_actions], "GithubActionsWithoutOidc"), + (vec![url], "Anonymous"), + (vec![("GITHUB_ACTIONS", "false")], "Anonymous"), + (vec![], "Anonymous"), + ] { + let auth = store_auth(&envs(&pairs)); + let kind = match auth { + StoreAuth::Anonymous => "Anonymous", + StoreAuth::GithubOidc(_) => "GithubOidc", + StoreAuth::GithubActionsWithoutOidc => "GithubActionsWithoutOidc", + }; + assert_eq!(kind, expected, "{pairs:?}"); + } + } + + /// Serve one request for each of `responses` on a loopback server, write + /// them in order, and close each connection. Returns the server's address + /// and the raw requests. + fn serve_each(responses: Vec<&'static [u8]>) -> (Str, std::thread::JoinHandle>>) { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = vt_str::format!("{}", listener.local_addr().unwrap()); + let server = std::thread::spawn(move || { + responses + .into_iter() + .map(|response| { + let (mut stream, _) = listener.accept().unwrap(); + let request = read_request(&mut stream); + stream.write_all(response).unwrap(); + request + }) + .collect() + }); + (address, server) + } + + /// Read a request with a `content-length` body. + fn read_request(stream: &mut std::net::TcpStream) -> Vec { + let mut request = Vec::new(); + let mut buf = [0; 4096]; + let head_end = loop { + let n = stream.read(&mut buf).unwrap(); + assert_ne!(n, 0, "connection closed before the request head ended"); + request.extend_from_slice(&buf[..n]); + if let Some(pos) = request.windows(4).position(|window| window == b"\r\n\r\n") { + break pos + 4; + } + }; + let content_length: usize = std::str::from_utf8(&request[..head_end]) + .unwrap() + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length").then(|| value.trim().parse().unwrap()) + }) + .unwrap_or(0); + while request.len() < head_end + content_length { + let n = stream.read(&mut buf).unwrap(); + assert_ne!(n, 0, "connection closed before the request body ended"); + request.extend_from_slice(&buf[..n]); + } + request + } + + async fn upload_to(clients: &RemoteClients, endpoint: &Arc) -> Result<(), UploadError> { + let key = cache_key(ResolvedGlobConfig::default_auto()); + let execution_key = ExecutionCacheKey::ExecAPI(Arc::from([])); + let value = CacheEntryValue { output_archive: None, ..cache_value() }; + let cache_dir = vt_path::current_dir().unwrap(); + clients + .upload(endpoint, &key, &execution_key, &value, &cache_dir, &CancellationToken::new()) + .await + } + + /// The message of `error` and each of its sources. + fn messages(error: &UploadError) -> Vec { + std::iter::successors(Some(error as &dyn std::error::Error), |err| err.source()) + .map(|err| vt_str::format!("{err}")) + .collect() + } + + #[tokio::test] + async fn uploads_stop_once_one_is_unauthorized() { + const TOKEN_FAILED: &[u8] = + b"HTTP/1.1 500 Internal Server Error\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"; + const UNAUTHORIZED: &[u8] = + b"HTTP/1.1 401 Unauthorized\r\ncontent-length: 7\r\nconnection: close\r\n\r\nno auth"; + const FORBIDDEN: &[u8] = + b"HTTP/1.1 403 Forbidden\r\ncontent-length: 6\r\nconnection: close\r\n\r\ndenied"; + for (github_actions, oidc, response, expected) in [ + ( + true, + true, + TOKEN_FAILED, + ["failed to get a GitHub Actions OIDC token", "HTTP status 500"].as_slice(), + ), + (false, false, UNAUTHORIZED, &["HTTP status 401", "no auth"]), + (false, false, FORBIDDEN, &["HTTP status 403", "denied"]), + ( + true, + false, + UNAUTHORIZED, + &["HTTP status 401", "grant `id-token: write` to this job", "no auth"], + ), + ] { + let (address, server) = serve_each(vec![response]); + let token_url = vt_str::format!("http://{address}/token?api-version=2.0"); + let mut pairs = vec![]; + if github_actions { + pairs.push(("GITHUB_ACTIONS", "true")); + } + if oidc { + pairs.push(("ACTIONS_ID_TOKEN_REQUEST_URL", token_url.as_str())); + pairs.push(("ACTIONS_ID_TOKEN_REQUEST_TOKEN", "request-token")); + } + let clients = RemoteClients::new(store_auth(&envs(&pairs))); + let endpoint = Arc::from(vt_str::format!("http://{address}/projects/test").as_str()); + + let first = upload_to(&clients, &endpoint).await.unwrap_err(); + assert!(matches!(first, UploadError::Unauthorized(_)), "{first:?}"); + assert_eq!(messages(&first), expected); + assert_eq!(server.join().unwrap().len(), 1); + + // The server is gone, so a request would be a network error. + for _ in 0..2 { + let skipped = upload_to(&clients, &endpoint).await.unwrap_err(); + assert!(matches!(skipped, UploadError::Unauthorized(_)), "{skipped:?}"); + assert_eq!(messages(&skipped), messages(&first)); + } + } + } + + #[tokio::test] + async fn other_upload_failures_do_not_stop_uploads() { + let (address, server) = serve_each(vec![ + b"HTTP/1.1 500 Internal Server Error\r\ncontent-length: 0\r\nconnection: close\r\n\r\n", + b"HTTP/1.1 200 OK\r\ncontent-length: 0\r\nconnection: close\r\n\r\n", + ]); + let clients = RemoteClients::new(StoreAuth::GithubActionsWithoutOidc); + let endpoint = Arc::from(vt_str::format!("http://{address}/projects/test").as_str()); + + let failed = upload_to(&clients, &endpoint).await.unwrap_err(); + assert!(matches!(failed, UploadError::Remote(_)), "{failed:?}"); + upload_to(&clients, &endpoint).await.unwrap(); + assert_eq!(server.join().unwrap().len(), 2); + } + #[tokio::test] async fn cancelling_stops_an_upload() { let (endpoint, requested) = serve_stalled(b""); diff --git a/crates/vt/src/session/mod.rs b/crates/vt/src/session/mod.rs index 6495080c7..3b0549a4c 100644 --- a/crates/vt/src/session/mod.rs +++ b/crates/vt/src/session/mod.rs @@ -597,13 +597,16 @@ impl<'a> Session<'a> { /// Lazily initializes and returns the execution cache. /// The cache is only created when first accessed to avoid `SQLite` race conditions - /// when multiple processes start simultaneously. + /// when multiple processes start simultaneously. Remote cache uploads + /// authenticate as the session envs allow. /// /// # Errors /// /// Returns an error if the cache database cannot be loaded or created. pub fn cache(&self) -> anyhow::Result<&ExecutionCache> { - self.cache.get_or_try_init(|| ExecutionCache::load_from_path(&self.cache_path)) + self.cache.get_or_try_init(|| { + ExecutionCache::load_from_path(&self.cache_path, cache::remote::store_auth(&self.envs)) + }) } pub fn workspace_path(&self) -> Arc { diff --git a/crates/vt_bin/src/vtt/main.rs b/crates/vt_bin/src/vtt/main.rs index e77c719df..beecff940 100644 --- a/crates/vt_bin/src/vtt/main.rs +++ b/crates/vt_bin/src/vtt/main.rs @@ -14,6 +14,7 @@ mod exit_on_ctrlc; mod grep_file; mod list_dir; mod mkdir; +mod oidc_remote_cache; mod pipe_stdin; mod print; mod print_color; @@ -36,7 +37,7 @@ fn main() { if args.len() < 2 { eprintln!("Usage: vtt [args...]"); eprintln!( - "Subcommands: barrier, check-tty, cp, exit, exit-on-ctrlc, grep-file, list-dir, mkdir, pipe-stdin, print, print-color, print-cwd, print-env, print-file, read-stdin, replace-file-content, rm, small_dev_shm, stalled-remote-cache, stat-file, stat-many, touch-file, write-file" + "Subcommands: barrier, check-tty, cp, exit, exit-on-ctrlc, grep-file, list-dir, mkdir, oidc-remote-cache, pipe-stdin, print, print-color, print-cwd, print-env, print-file, read-stdin, replace-file-content, rm, small_dev_shm, stalled-remote-cache, stat-file, stat-many, touch-file, write-file" ); std::process::exit(1); } @@ -56,6 +57,7 @@ fn main() { } "list-dir" => list_dir::run(&args[2..]), "mkdir" => mkdir::run(&args[2..]), + "oidc-remote-cache" => oidc_remote_cache::run(&args[2..]), "pipe-stdin" => pipe_stdin::run(&args[2..]), "print" => { print::run(&args[2..]); diff --git a/crates/vt_bin/src/vtt/oidc_remote_cache.rs b/crates/vt_bin/src/vtt/oidc_remote_cache.rs new file mode 100644 index 000000000..da41fa8bc --- /dev/null +++ b/crates/vt_bin/src/vtt/oidc_remote_cache.rs @@ -0,0 +1,183 @@ +use std::{ + io::{Read as _, Write as _}, + net::{TcpListener, TcpStream}, + sync::{Arc, Mutex}, +}; + +const BASE_PATH: &str = "/projects/test"; +const REQUEST_TOKEN: &str = "request-token"; +/// A JWT with a fake signature that expires in 2100. +const TOKEN: &str = "eyJhbGciOiJSUzI1NiIsInR5cCI6IkpXVCJ9.eyJleHAiOjQxMDI0NDQ4MDB9.c2lnbmF0dXJl"; + +/// oidc-remote-cache \[`--without-oidc`\] `` \[``...\] +/// +/// Runs `` with `VP_REMOTE_CACHE_URL` set to a loopback remote cache +/// that also serves a fake GitHub Actions OIDC token endpoint, and with +/// `ACTIONS_ID_TOKEN_REQUEST_URL` and `ACTIONS_ID_TOKEN_REQUEST_TOKEN` set to +/// request tokens from it, unless `--without-oidc` is given. Then exits with +/// the command's exit code, after printing a line for each request to stderr. +/// +/// - `GET /token` responds with a token if the request carries the request +/// token and asks for the endpoint as the audience, and with 401 or 400 +/// otherwise. +/// - `POST {BASE_PATH}/fetch` responds with 404, or 400 if it carries +/// credentials. +/// - `POST {BASE_PATH}/store` responds with 200 if it carries the token, 401 +/// if it doesn't, and 400 if it contains the request token. +pub fn run(args: &[String]) -> Result<(), Box> { + let (with_oidc, args) = match args { + [flag, rest @ ..] if flag == "--without-oidc" => (false, rest), + _ => (true, args), + }; + let [program, args @ ..] = args else { + return Err("Usage: vtt oidc-remote-cache [--without-oidc] [args...]".into()); + }; + + let listener = TcpListener::bind("127.0.0.1:0")?; + let origin = format!("http://{}", listener.local_addr()?); + let endpoint = format!("{origin}{BASE_PATH}"); + let log = Arc::new(Mutex::new(Vec::new())); + { + let endpoint = endpoint.clone(); + let log = Arc::clone(&log); + std::thread::spawn(move || { + for stream in listener.incoming().filter_map(Result::ok) { + handle(stream, &endpoint, &log); + } + }); + } + + let mut command = std::process::Command::new(program); + command.args(args).env("VP_REMOTE_CACHE_URL", &endpoint); + if with_oidc { + command + .env("ACTIONS_ID_TOKEN_REQUEST_URL", format!("{origin}/token?api-version=2.0")) + .env("ACTIONS_ID_TOKEN_REQUEST_TOKEN", REQUEST_TOKEN); + } + let status = command.status()?; + for line in log.lock().unwrap().iter() { + eprintln!("{line}"); + } + std::process::exit(status.code().unwrap_or(1)); +} + +/// Respond to one request and close the connection. The request is logged +/// before the response is written, so it's in `log` once the client has the +/// response. +fn handle(mut stream: TcpStream, endpoint: &str, log: &Mutex>) -> Option<()> { + let request = read_request(&mut stream)?; + let head = String::from_utf8_lossy(&request.bytes[..request.head_len]).into_owned(); + let mut lines = head.split("\r\n"); + let mut request_line = lines.next()?.split(' '); + let (method, target) = (request_line.next()?, request_line.next()?); + let authorization = lines.find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("authorization").then(|| value.trim().to_owned()) + }); + let (path, query) = target.split_once('?').unwrap_or((target, "")); + let leaks_request_token = + request.bytes.windows(REQUEST_TOKEN.len()).any(|window| window == REQUEST_TOKEN.as_bytes()); + + let (prefix, route, status) = match (method, path.strip_prefix(BASE_PATH)) { + ("GET", _) if path == "/token" => { + let status = if authorization.as_deref() != Some(&format!("Bearer {REQUEST_TOKEN}")) { + "401 Unauthorized" + } else if query_value(query, "audience").as_deref() != Some(endpoint) { + "400 Bad Request" + } else { + "200 OK" + }; + ("github-oidc", path, status) + } + ("POST", Some(route @ "/fetch")) => ( + "remote-cache", + route, + if authorization.is_some() { "400 Bad Request" } else { "404 Not Found" }, + ), + ("POST", Some(route @ "/store")) => { + let status = if leaks_request_token { + "400 Bad Request" + } else if authorization.as_deref() == Some(&format!("Bearer {TOKEN}")) { + "200 OK" + } else { + "401 Unauthorized" + }; + ("remote-cache", route, status) + } + _ => ("remote-cache", path, "404 Not Found"), + }; + let body = if prefix == "github-oidc" && status == "200 OK" { + format!(r#"{{"value":"{TOKEN}"}}"#) + } else { + String::new() + }; + let code = status.split(' ').next()?; + log.lock().unwrap().push(format!("[{prefix}] {method} {route} {code}")); + let response = format!( + "HTTP/1.1 {status}\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}", + body.len() + ); + stream.write_all(response.as_bytes()).ok() +} + +struct Request { + bytes: Vec, + head_len: usize, +} + +/// Read a request with a `content-length` body. `None` if the connection +/// closes first. +fn read_request(stream: &mut TcpStream) -> Option { + let mut bytes = Vec::new(); + let mut buf = [0; 4096]; + let head_len = loop { + let n = stream.read(&mut buf).ok().filter(|&n| n > 0)?; + bytes.extend_from_slice(&buf[..n]); + if let Some(pos) = bytes.windows(4).position(|window| window == b"\r\n\r\n") { + break pos + 4; + } + }; + let content_length: usize = String::from_utf8_lossy(&bytes[..head_len]) + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length").then(|| value.trim().parse().ok())? + }) + .unwrap_or(0); + while bytes.len() < head_len + content_length { + let n = stream.read(&mut buf).ok().filter(|&n| n > 0)?; + bytes.extend_from_slice(&buf[..n]); + } + Some(Request { bytes, head_len }) +} + +/// The decoded value of the first `name` parameter in a URL-encoded query. +fn query_value(query: &str, name: &str) -> Option { + query.split('&').find_map(|pair| { + let (key, value) = pair.split_once('=')?; + (key == name).then(|| percent_decode(value))? + }) +} + +fn percent_decode(value: &str) -> Option { + let mut bytes = Vec::with_capacity(value.len()); + let mut rest = value.as_bytes(); + while let [byte, tail @ ..] = rest { + match byte { + b'%' => { + let hex = std::str::from_utf8(tail.get(..2)?).ok()?; + bytes.push(u8::from_str_radix(hex, 16).ok()?); + rest = &tail[2..]; + } + b'+' => { + bytes.push(b' '); + rest = tail; + } + _ => { + bytes.push(*byte); + rest = tail; + } + } + } + String::from_utf8(bytes).ok() +} diff --git a/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots.toml b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots.toml index ba61d0b67..93939e057 100644 --- a/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots.toml +++ b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots.toml @@ -379,3 +379,47 @@ steps = [ "fail-during-build", ], comment = "The endpoint never responds. fail exits while build's fetch is in flight, which stops the fetch, and build doesn't start." }, ] + +[[e2e]] +name = "github_oidc_upload" +steps = [ + { argv = [ + "vtt", + "oidc-remote-cache", + "vt", + "run", + "verify", + ], envs = [ + [ + "VP_REMOTE_CACHE", + "read-write", + ], + ], comment = "The first upload requests an OIDC token for the endpoint, and the second reuses it. Fetches send no token." }, +] + +[[e2e]] +name = "github_oidc_missing_permission" +steps = [ + { argv = [ + "vtt", + "oidc-remote-cache", + "--without-oidc", + "vt", + "run", + "verify", + ], envs = [ + [ + "VP_REMOTE_CACHE", + "read-write", + ], + [ + "GITHUB_ACTIONS", + "true", + ], + ], comment = "Without the OIDC variables, the first upload has no token and is rejected. The second upload is skipped without a request, with the same reason. The tasks succeed." }, + { argv = [ + "vt", + "run", + "--last-details", + ], comment = "The details suggest granting `id-token: write`." }, +] diff --git a/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_missing_permission.md b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_missing_permission.md new file mode 100644 index 000000000..8ca2a4327 --- /dev/null +++ b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_missing_permission.md @@ -0,0 +1,45 @@ +# github_oidc_missing_permission + +## `VP_REMOTE_CACHE=read-write GITHUB_ACTIONS=true vtt oidc-remote-cache --without-oidc vt run verify` + +Without the OIDC variables, the first upload has no token and is rejected. The second upload is skipped without a request, with the same reason. The tasks succeed. + +``` +$ vtt write-file dist/output.txt built + +$ vtt print verified +verified + +--- +vt run: 0/2 cache hit (0%). remote-cache#build (and 1 more) not uploaded to the remote cache: HTTP status 401. (Run `vt run --last-details` for full details) +[remote-cache] POST /fetch 404 +[remote-cache] POST /store 401 +[remote-cache] POST /fetch 404 +``` + +## `vt run --last-details` + +The details suggest granting `id-token: write`. + +``` + +━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ + Vite+ Task Runner • Execution Summary +━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ + +Statistics: 2 tasks • 0 cache hits • 2 cache misses +Performance: 0% cache hit rate + +Task Details: +──────────────────────────────────────────────── + [1] remote-cache#build: $ vtt write-file dist/output.txt built ✓ + → Cache miss: no previous cache entry found + ⚠ Not uploaded to the remote cache: HTTP status 401 + ↳ grant `id-token: write` to this job + ······················································· + [2] remote-cache#verify: $ vtt print verified ✓ + → Cache miss: no previous cache entry found + ⚠ Not uploaded to the remote cache: HTTP status 401 + ↳ grant `id-token: write` to this job +━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ +``` diff --git a/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_upload.md b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_upload.md new file mode 100644 index 000000000..629289abd --- /dev/null +++ b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_upload.md @@ -0,0 +1,20 @@ +# github_oidc_upload + +## `VP_REMOTE_CACHE=read-write vtt oidc-remote-cache vt run verify` + +The first upload requests an OIDC token for the endpoint, and the second reuses it. Fetches send no token. + +``` +$ vtt write-file dist/output.txt built + +$ vtt print verified +verified + +--- +vt run: 0/2 cache hit (0%). (Run `vt run --last-details` for full details) +[remote-cache] POST /fetch 404 +[github-oidc] GET /token 200 +[remote-cache] POST /store 200 +[remote-cache] POST /fetch 404 +[remote-cache] POST /store 200 +``` diff --git a/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/vite-task.json b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/vite-task.json index 14be7f2eb..b45fc17b9 100644 --- a/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/vite-task.json +++ b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/vite-task.json @@ -17,6 +17,14 @@ "all": { "command": "vt run build && vt run check" }, + "verify": { + "command": "vtt print verified", + "dependsOn": ["build"], + "cache": { + "input": ["src/**"], + "output": [] + } + }, "fail": { "command": "vtt exit 1", "cache": false diff --git a/crates/vt_graph/run-config.ts b/crates/vt_graph/run-config.ts index b1d68ee16..bd9cacfb4 100644 --- a/crates/vt_graph/run-config.ts +++ b/crates/vt_graph/run-config.ts @@ -29,6 +29,9 @@ export type InputBase = "package" | "workspace"; export type RemoteCacheConfig = { /** * HTTP or HTTPS namespace endpoint. Overridden by `VP_REMOTE_CACHE_URL`. + * + * OIDC tokens for uploads are requested with this endpoint as their + * audience, without its query, userinfo, or trailing slash. */ url: string, }; @@ -131,6 +134,10 @@ scripts?: boolean, tasks?: boolean, /** * Remote cache shared by tasks in the workspace. + * + * In `read-write` mode, uploads from a GitHub Actions job granted + * `id-token: write` authenticate with an OIDC token for the endpoint. + * Other uploads send no credentials. */ remote?: RemoteCacheConfig, }; diff --git a/crates/vt_graph/src/config/user.rs b/crates/vt_graph/src/config/user.rs index 0564a97eb..e297fd2a3 100644 --- a/crates/vt_graph/src/config/user.rs +++ b/crates/vt_graph/src/config/user.rs @@ -368,6 +368,10 @@ pub enum UserGlobalCacheConfig { tasks: Option, /// Remote cache shared by tasks in the workspace. + /// + /// In `read-write` mode, uploads from a GitHub Actions job granted + /// `id-token: write` authenticate with an OIDC token for the endpoint. + /// Other uploads send no credentials. remote: Option, }, } @@ -412,6 +416,9 @@ impl ResolvedGlobalCacheConfig { #[serde(deny_unknown_fields, rename_all = "camelCase")] pub struct UserRemoteCacheConfig { /// HTTP or HTTPS namespace endpoint. Overridden by `VP_REMOTE_CACHE_URL`. + /// + /// OIDC tokens for uploads are requested with this endpoint as their + /// audience, without its query, userinfo, or trailing slash. pub url: Arc, } diff --git a/crates/vt_remote_cache/Cargo.toml b/crates/vt_remote_cache/Cargo.toml index 750b76f06..f4ccacd5d 100644 --- a/crates/vt_remote_cache/Cargo.toml +++ b/crates/vt_remote_cache/Cargo.toml @@ -8,6 +8,7 @@ publish = false rust-version.workspace = true [dependencies] +base64 = { workspace = true } bytes = { workspace = true } ciborium = { workspace = true } reqwest = { workspace = true, features = [ @@ -19,8 +20,9 @@ reqwest = { workspace = true, features = [ rustls = { workspace = true, features = ["ring", "std"] } serde = { workspace = true, features = ["derive"] } serde_bytes = { workspace = true } +serde_json = { workspace = true } thiserror = { workspace = true } -tokio = { workspace = true, features = ["fs"] } +tokio = { workspace = true, features = ["fs", "sync"] } url = { workspace = true } vt_path = { workspace = true } vt_str = { workspace = true } diff --git a/crates/vt_remote_cache/README.md b/crates/vt_remote_cache/README.md index 150cc62fb..2e3a6f572 100644 --- a/crates/vt_remote_cache/README.md +++ b/crates/vt_remote_cache/README.md @@ -2,7 +2,7 @@ Client for the [remote cache server API](https://github.com/voidzero-dev/vite-task/pull/713). It treats keys, values, and blobs as opaque bytes. Encoding task cache entries into them is up to the caller. -`Client::new` takes the configured endpoint, which can include a namespace path, such as `https://cache.example.com/projects/my-project`. Each operation appends its route to that path, so a store goes to `https://cache.example.com/projects/my-project/store`. Endpoints that aren't HTTP or HTTPS URLs are rejected. +`Client::new` takes the configured endpoint, which can include a namespace path, such as `https://cache.example.com/projects/my-project`. Each operation appends its route to that path, so a store goes to `https://cache.example.com/projects/my-project/store`. Endpoints that aren't HTTP or HTTPS URLs are rejected. It also takes a `StoreAuth`, which says how stores authenticate. Fetches and downloads never send credentials. `Client::fetch` sends a key and a secondary key as a CBOR map of byte strings. A 404 response means neither key matched, so `fetch` returns `None` without decoding the body. The body of a 200 response is a CBOR map tagged by `kind`, and `Fetched` decodes only two kinds: `exact`, an exact match with its value and blob ID, and `fallback`, a fallback match with the key it's stored under. Any other kind is a malformed response. An exact match must include `blob_id`, which is null when there's no blob. The fallback's value and blob ID aren't decoded. @@ -14,4 +14,12 @@ Only HTTP 200 counts as success for every operation, except that a fetch also ac reqwest configures TLS. It uses the process's default rustls crypto provider, which the client installs as ring unless one is already installed, and verifies certificates with the operating system's verifier. Requests go through the proxy set in `HTTPS_PROXY`, `HTTP_PROXY`, or `ALL_PROXY`, except for hosts in `NO_PROXY`. Without those variables, the system proxy settings are used on macOS and Windows. Connections time out after 10 seconds. Reads time out after 60 seconds, and until the response headers arrive, that limit also covers sending the request. -`Error` names the kind of failure: an invalid endpoint, a client that couldn't be created, a blob file that couldn't be read, a network error (including timeouts and responses that end early), a status other than 200 (or 404, for a fetch), or a malformed fetch response. Its messages contain no OS-specific details, so they can be shown to users as is. The details are in the source: the underlying error, the parse error for an endpoint that isn't a URL, or the message in an error response's body. Network errors leave out the request URL, since the endpoint may contain credentials, such as a token in its query. +`Error` names the kind of failure: an invalid endpoint, a client that couldn't be created, a blob file that couldn't be read, a network error (including timeouts and responses that end early), a status other than 200 (or 404, for a fetch), a malformed fetch response, or a failed OIDC token request. Its messages contain no OS-specific details, so they can be shown to users as is. The details are in the source: the underlying error, the parse error for an endpoint that isn't a URL, or the message in an error response's body. Network errors leave out the request URL, since the endpoint may contain credentials, such as a token in its query. + +## GitHub Actions OIDC + +With `StoreAuth::GithubOidc`, which holds the values of a GitHub Actions job's `ACTIONS_ID_TOKEN_REQUEST_URL` and `ACTIONS_ID_TOKEN_REQUEST_TOKEN`, each store sends a GitHub OIDC token as `Authorization: Bearer `. GitHub sets these variables only in jobs granted `id-token: write`. Before the first store, the client requests a token with `GET &audience=`, adding `audience` to the URL's query, and `Authorization: Bearer `. The request token goes only to the token endpoint. The audience is the endpoint without its query, fragment, userinfo, or trailing slash, such as `https://cache.example.com/projects/my-project`. The token endpoint must respond with 200, without a redirect, and a JSON object whose `value` is the token, a JWT. + +The client keeps the token in memory and reuses it for later stores until it expires within 60 seconds, reading its `exp` claim without verifying it. Stores that need a new token at the same time wait for a single request. If a token request fails, the client keeps its error and returns it for every later store without requesting again. The error leaves out the token endpoint's URL and response body. If the endpoint has userinfo, the token replaces its basic credentials. `Debug` output leaves out the tokens and the endpoint. + +`StoreAuth::GithubActionsWithoutOidc` is for GitHub Actions jobs that lack these variables. Stores send no credentials, like `StoreAuth::Anonymous`, but a 401 response is `Error::MissingIdTokenPermission`, whose message is the same as for any other 401 and whose cause suggests granting `id-token: write`. `Error::is_unauthorized` is true for it, a failed token request, and 401 and 403 responses. diff --git a/crates/vt_remote_cache/src/lib.rs b/crates/vt_remote_cache/src/lib.rs index c5f91b2f9..f9f4ff52c 100644 --- a/crates/vt_remote_cache/src/lib.rs +++ b/crates/vt_remote_cache/src/lib.rs @@ -1,12 +1,17 @@ //! Client for the remote cache server API. Keys, values, and blobs are opaque //! bytes; the caller decides what they contain. -use std::time::Duration; +use std::{ + fmt, + sync::Arc, + time::{Duration, SystemTime}, +}; +use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; use bytes::Bytes; use reqwest::{ Response, StatusCode, - header::CONTENT_TYPE, + header::{AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderValue}, multipart::{Form, Part}, }; use serde::{Deserialize, Serialize}; @@ -21,6 +26,10 @@ const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); /// arrive, it also bounds sending the request. const READ_TIMEOUT: Duration = Duration::from_secs(60); +/// A cached OIDC token is replaced once it expires within this long, so it +/// doesn't expire while a store is being sent. +const TOKEN_REFRESH_MARGIN: Duration = Duration::from_secs(60); + /// A failed remote cache operation. The messages name only the kind of /// failure, so they are the same on every platform. The details, if any, are /// in the source. @@ -51,6 +60,29 @@ pub enum Error { /// The body of a 200 fetch response isn't an exact or fallback match. #[error("malformed response")] MalformedResponse(#[source] ciborium::de::Error), + /// No GitHub Actions OIDC token could be obtained for a store. Later + /// stores return the same error without requesting another token. + #[error("failed to get a GitHub Actions OIDC token")] + OidcToken(#[source] Arc), + /// The server responded to a store with 401 from a GitHub Actions job + /// that can't request OIDC tokens, so the store had none. The message is + /// the same as for any other 401 response. + #[error("HTTP status 401")] + MissingIdTokenPermission(#[source] IdTokenPermissionHint), +} + +impl Error { + /// Whether the error means the store wasn't authorized: no OIDC token + /// could be obtained, or the server responded with 401 or 403. + #[must_use] + pub const fn is_unauthorized(&self) -> bool { + matches!( + self, + Self::OidcToken(_) + | Self::MissingIdTokenPermission(_) + | Self::Status(StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN, _) + ) + } } /// The message in the body of an error response. @@ -58,6 +90,136 @@ pub enum Error { #[error("{0}")] pub struct ServerMessage(Str); +/// Why no GitHub Actions OIDC token could be obtained. Neither the messages +/// nor the sources contain the token endpoint's URL, the request token, or +/// the response body. +#[derive(Debug, thiserror::Error)] +pub enum OidcTokenError { + /// The request URL isn't a URL, or the request token can't be sent in a + /// header. + #[error("invalid token request")] + InvalidRequest, + /// No complete response arrived. + #[error("network error")] + Network(#[source] reqwest::Error), + /// The token endpoint responded with a status other than 200. Redirects + /// aren't followed, so a redirect is a status like any other. + #[error("HTTP status {}", .0.as_u16())] + Status(StatusCode), + /// The body of the 200 response isn't a JSON object whose `value` is a + /// JWT with an integer `exp` claim. + #[error("malformed response")] + MalformedResponse, +} + +/// The cause of a 401 response to a store sent without a token from a GitHub +/// Actions job: the job can't request OIDC tokens. The source is the message +/// in the response body, if any. +#[derive(Debug, thiserror::Error)] +#[error("grant `id-token: write` to this job")] +pub struct IdTokenPermissionHint(#[source] Option); + +/// How [`Client::store`] authenticates. Fetches and downloads never send +/// credentials. +#[derive(Debug, Clone, Default)] +pub enum StoreAuth { + /// Send no credentials. + #[default] + Anonymous, + /// Send a GitHub Actions OIDC token as `Authorization: Bearer `. + GithubOidc(GithubOidc), + /// Send no credentials from a GitHub Actions job that can't request OIDC + /// tokens. A 401 response is [`Error::MissingIdTokenPermission`]. + GithubActionsWithoutOidc, +} + +/// What a GitHub Actions job uses to request OIDC tokens. +/// +/// These are the values of `ACTIONS_ID_TOKEN_REQUEST_URL` and +/// `ACTIONS_ID_TOKEN_REQUEST_TOKEN`. GitHub sets them only in jobs granted +/// `id-token: write`. `Debug` leaves both out. +#[derive(Clone)] +pub struct GithubOidc { + request_url: Arc, + request_token: Arc, +} + +impl GithubOidc { + #[must_use] + pub fn new(request_url: &str, request_token: &str) -> Self { + Self { request_url: Arc::from(request_url), request_token: Arc::from(request_token) } + } +} + +impl fmt::Debug for GithubOidc { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("GithubOidc").finish_non_exhaustive() + } +} + +/// The OIDC token that stores share. A failed request is kept, so later +/// stores get its error without another request. +enum TokenState { + Empty, + /// `authorization` is `Bearer `, and `expires_at` is the token's + /// `exp` claim, in seconds since the Unix epoch. + Valid { + authorization: HeaderValue, + expires_at: u64, + }, + Failed(Arc), +} + +/// The response of the token endpoint. +#[derive(Deserialize)] +struct TokenResponse { + value: Box, +} + +/// The only claim read from a token. +#[derive(Deserialize)] +struct TokenClaims { + exp: u64, +} + +/// The audience of tokens for `endpoint`: the endpoint without its query, +/// fragment, userinfo, or trailing slash. +fn audience(endpoint: &Url) -> Str { + let mut url = endpoint.clone(); + url.set_query(None); + url.set_fragment(None); + // These fail only for URLs that can't have userinfo, which have none. + let _ = url.set_username(""); + let _ = url.set_password(None); + let url = url.as_str(); + Str::from(url.strip_suffix('/').unwrap_or(url)) +} + +/// The `exp` claim in the payload of `jwt`, read without verifying the +/// signature. +fn jwt_expiry(jwt: &str) -> Option { + let payload = URL_SAFE_NO_PAD.decode(jwt.split('.').nth(1)?).ok()?; + let TokenClaims { exp } = serde_json::from_slice(&payload).ok()?; + Some(exp) +} + +/// Whether a token that expires at `expires_at`, in seconds since the Unix +/// epoch, expires within [`TOKEN_REFRESH_MARGIN`]. +fn expires_soon(expires_at: u64) -> bool { + let now = SystemTime::now() + .duration_since(SystemTime::UNIX_EPOCH) + .map_or(0, |since_epoch| since_epoch.as_secs()); + expires_at <= now.saturating_add(TOKEN_REFRESH_MARGIN.as_secs()) +} + +/// The `Authorization` header value `Bearer `, marked sensitive so +/// it's left out of debug output. `None` if `token` can't be sent in a header. +fn bearer(token: &str) -> Option { + let mut value = HeaderValue::from_bytes(&[b"Bearer ", token.as_bytes()].concat()).ok()?; + value.set_sensitive(true); + Some(value) +} + /// A blob being downloaded. #[derive(Debug)] pub struct Download { @@ -121,24 +283,36 @@ pub enum Fetched { } /// A client for one remote cache endpoint. -#[derive(Debug)] pub struct Client { http: reqwest::Client, fetch_url: Url, store_url: Url, /// `{endpoint}/blob`, to which each download appends a blob ID. blob_url: Url, + store_auth: StoreAuth, + /// The audience OIDC tokens are requested for. See [`audience`]. + audience: Str, + oidc_token: tokio::sync::Mutex, +} + +impl fmt::Debug for Client { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + // Leave out the cached token and the endpoint, which may contain + // credentials. + f.debug_struct("Client").finish_non_exhaustive() + } } impl Client { /// Create a client for `endpoint`, a base URL that may include a /// namespace path, such as `https://cache.example.com/projects/my-project`. + /// Stores authenticate with `store_auth`. /// /// # Errors /// /// Returns [`Error::InvalidEndpoint`] if `endpoint` isn't a usable URL, or /// [`Error::HttpClient`] if the HTTP client can't be created. - pub fn new(endpoint: &str) -> Result { + pub fn new(endpoint: &str, store_auth: StoreAuth) -> Result { let endpoint = parse_endpoint(endpoint)?; let fetch_url = route_url(&endpoint, "fetch")?; let store_url = route_url(&endpoint, "store")?; @@ -155,7 +329,15 @@ impl Client { .redirect(reqwest::redirect::Policy::none()) .build() .map_err(Error::HttpClient)?; - Ok(Self { http, fetch_url, store_url, blob_url }) + Ok(Self { + http, + fetch_url, + store_url, + blob_url, + store_auth, + audience: audience(&endpoint), + oidc_token: tokio::sync::Mutex::new(TokenState::Empty), + }) } /// Fetch the entry stored under `key` with `POST {endpoint}/fetch`, @@ -204,12 +386,14 @@ impl Client { /// Store `value` under `key` with `POST {endpoint}/store`, uploading the /// file at `blob` as its blob. Fetches that match no key fall back to this - /// entry through `secondary_key`. + /// entry through `secondary_key`. The request authenticates as the client's + /// [`StoreAuth`] says, first requesting an OIDC token if it needs one. /// /// # Errors /// - /// Returns an error if the blob file can't be opened, the request fails, - /// or the server responds with a status other than 200. + /// Returns an error if no OIDC token can be obtained, the blob file can't + /// be opened, the request fails, or the server responds with a status + /// other than 200. pub async fn store( &self, key: &[u8], @@ -217,22 +401,93 @@ impl Client { value: &[u8], blob: Option<&AbsolutePath>, ) -> Result<(), Error> { + let authorization = match &self.store_auth { + StoreAuth::GithubOidc(oidc) => { + Some(self.oidc_authorization(oidc).await.map_err(Error::OidcToken)?) + } + StoreAuth::Anonymous | StoreAuth::GithubActionsWithoutOidc => None, + }; let metadata = StoreMetadata { key, secondary_key, value }; let mut form = Form::new().part("metadata", metadata_part(&metadata)); if let Some(blob) = blob { form = form.part("blob", blob_part(blob).await?); } + let mut request = self.http.post(self.store_url.clone()).multipart(form); + if let Some(authorization) = authorization { + // `headers` replaces the basic credentials taken from userinfo in + // the endpoint, so only the token is sent. + request = request.headers(HeaderMap::from_iter([(AUTHORIZATION, authorization)])); + } + let response = request.send().await.map_err(network_error)?; + let response = match check_status(response).await { + Err(Error::Status(StatusCode::UNAUTHORIZED, message)) + if matches!(self.store_auth, StoreAuth::GithubActionsWithoutOidc) => + { + return Err(Error::MissingIdTokenPermission(IdTokenPermissionHint(message))); + } + response => response?, + }; + // The response's blob ID isn't needed. Read the body anyway, so the + // connection can be reused. + response.bytes().await.map_err(network_error)?; + Ok(()) + } + + /// The `Authorization` header value for a store: the cached token, or a + /// new one if it's missing or about to expire. The lock is held during + /// the request, so stores that need a token at the same time share it. A + /// new token is used even if it's about to expire itself; the server + /// decides whether it's valid. + async fn oidc_authorization( + &self, + oidc: &GithubOidc, + ) -> Result> { + let mut state = self.oidc_token.lock().await; + match &*state { + TokenState::Valid { authorization, expires_at } if !expires_soon(*expires_at) => { + return Ok(authorization.clone()); + } + TokenState::Failed(err) => return Err(Arc::clone(err)), + TokenState::Valid { .. } | TokenState::Empty => {} + } + let requested = self.request_oidc_token(oidc).await.map_err(Arc::new); + *state = match &requested { + Ok((authorization, expires_at)) => { + TokenState::Valid { authorization: authorization.clone(), expires_at: *expires_at } + } + Err(err) => TokenState::Failed(Arc::clone(err)), + }; + requested.map(|(authorization, _)| authorization) + } + + /// Request a token with `GET &audience=`. Returns + /// its `Authorization` header value and its `exp` claim. + async fn request_oidc_token( + &self, + oidc: &GithubOidc, + ) -> Result<(HeaderValue, u64), OidcTokenError> { + let mut url = Url::parse(&oidc.request_url).map_err(|_| OidcTokenError::InvalidRequest)?; + url.query_pairs_mut().append_pair("audience", &self.audience); + let request_authorization = + bearer(&oidc.request_token).ok_or(OidcTokenError::InvalidRequest)?; let response = self .http - .post(self.store_url.clone()) - .multipart(form) + .get(url) + .header(AUTHORIZATION, request_authorization) .send() .await - .map_err(network_error)?; - // The response's blob ID isn't needed. Read the body anyway, so the - // connection can be reused. - check_status(response).await?.bytes().await.map_err(network_error)?; - Ok(()) + .map_err(|err| OidcTokenError::Network(err.without_url()))?; + if response.status() != StatusCode::OK { + return Err(OidcTokenError::Status(response.status())); + } + let body = + response.bytes().await.map_err(|err| OidcTokenError::Network(err.without_url()))?; + // The parse error isn't kept, because it can quote the body. + let TokenResponse { value } = + serde_json::from_slice(&body).map_err(|_| OidcTokenError::MalformedResponse)?; + let expires_at = jwt_expiry(&value).ok_or(OidcTokenError::MalformedResponse)?; + let authorization = bearer(&value).ok_or(OidcTokenError::MalformedResponse)?; + Ok((authorization, expires_at)) } } @@ -322,7 +577,31 @@ mod tests { fn rejects_endpoints_that_are_not_http_urls() { for endpoint in ["cache.example/projects/test", "ftp://cache.example", "mailto:a@b.example"] { - assert!(matches!(Client::new(endpoint), Err(Error::InvalidEndpoint(_))), "{endpoint}"); + assert!( + matches!( + Client::new(endpoint, StoreAuth::Anonymous), + Err(Error::InvalidEndpoint(_)) + ), + "{endpoint}" + ); + } + } + + #[test] + fn audience_is_the_endpoint_without_query_userinfo_or_trailing_slash() { + for (endpoint, expected) in [ + ("https://cache.example/projects/test", "https://cache.example/projects/test"), + ("https://cache.example/projects/test/", "https://cache.example/projects/test"), + ("https://cache.example", "https://cache.example"), + ("https://cache.example/ns?token=a#top", "https://cache.example/ns"), + ("https://user:password@cache.example:8443/ns", "https://cache.example:8443/ns"), + ("HTTPS://Cache.Example:443/ns", "https://cache.example/ns"), + ] { + assert_eq!( + audience(&parse_endpoint(endpoint).unwrap()).as_str(), + expected, + "{endpoint}" + ); } } @@ -405,11 +684,31 @@ mod tests { haystack.windows(needle.len()).any(|window| window == needle) } + /// A response with `status_line` and `body` that closes the connection, + /// so the client opens a new one for its next request. + fn http_response(status_line: &str, body: &[u8]) -> Vec { + let head = vt_str::format!( + "{status_line}\r\ncontent-length: {}\r\nconnection: close\r\n\r\n", + body.len() + ); + [head.as_bytes(), body].concat() + } + /// Accept one HTTP request, respond with `status_line` and `body`, and /// return the raw request. fn serve_once(listener: &TcpListener, status_line: &str, body: &[u8]) -> Vec { - let headers = vt_str::format!("{status_line}\r\ncontent-length: {}\r\n\r\n", body.len()); - serve_raw_once(listener, &[headers.as_bytes(), body].concat()) + serve_raw_once(listener, &http_response(status_line, body)) + } + + /// Accept one request for each of `responses`, write them in order, and + /// return the raw requests. + fn serve_each( + listener: TcpListener, + responses: Vec>, + ) -> std::thread::JoinHandle>> { + std::thread::spawn(move || { + responses.iter().map(|response| serve_raw_once(&listener, response)).collect() + }) } /// Accept one HTTP request, write `response`, close the connection, and @@ -444,8 +743,64 @@ mod tests { } fn client_for(listener: &TcpListener) -> Client { + client_with_auth(listener, StoreAuth::Anonymous) + } + + fn client_with_auth(listener: &TcpListener, store_auth: StoreAuth) -> Client { + let port = listener.local_addr().unwrap().port(); + Client::new(&vt_str::format!("http://127.0.0.1:{port}/projects/test"), store_auth).unwrap() + } + + const REQUEST_TOKEN: &str = "request-token-secret"; + + /// A client whose token endpoint is `/token` on the same server. + fn oidc_client_for(listener: &TcpListener) -> Client { let port = listener.local_addr().unwrap().port(); - Client::new(&vt_str::format!("http://127.0.0.1:{port}/projects/test")).unwrap() + let request_url = vt_str::format!("http://127.0.0.1:{port}/token?api-version=2.0"); + client_with_auth( + listener, + StoreAuth::GithubOidc(GithubOidc::new(&request_url, REQUEST_TOKEN)), + ) + } + + fn unix_now() -> u64 { + SystemTime::now().duration_since(SystemTime::UNIX_EPOCH).unwrap().as_secs() + } + + /// A JWT with a fake signature that expires at `exp`, in seconds since + /// the Unix epoch. + fn jwt(exp: u64) -> Str { + let header = URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256","typ":"JWT"}"#); + let payload = URL_SAFE_NO_PAD.encode(vt_str::format!(r#"{{"exp":{exp}}}"#).as_str()); + vt_str::format!("{header}.{payload}.c2lnbmF0dXJl") + } + + fn token_response(token: &str) -> Vec { + http_response("HTTP/1.1 200 OK", vt_str::format!(r#"{{"value":"{token}"}}"#).as_bytes()) + } + + fn stored() -> Vec { + http_response("HTTP/1.1 200 OK", b"") + } + + fn authorization_line(token: &str) -> Vec { + vt_str::format!("authorization: Bearer {token}\r\n").as_bytes().to_vec() + } + + fn is_token_request(request: &[u8]) -> bool { + request.starts_with(b"GET /token?") + } + + fn is_store_request(request: &[u8]) -> bool { + request.starts_with(b"POST /projects/test/store HTTP/1.1\r\n") + } + + /// The message of `error` and each of its sources, and its debug output. + fn error_texts(error: &Error) -> Vec { + std::iter::successors(Some(error as &dyn std::error::Error), |err| err.source()) + .map(|err| vt_str::format!("{err}")) + .chain([vt_str::format!("{error:?}")]) + .collect() } /// The `metadata` part of a store request for key `k`, secondary key `s`, @@ -611,6 +966,7 @@ mod tests { assert!(request.starts_with(b"POST /projects/test/store HTTP/1.1\r\n")); assert!(contains(&request, &metadata_part_bytes())); assert!(!contains(&request, b"name=\"blob\"")); + assert!(!contains(&request, b"authorization:")); } #[tokio::test] @@ -631,16 +987,202 @@ mod tests { #[tokio::test] async fn network_errors_leave_out_the_url() { - let client = - Client::new("http://user:password@127.0.0.1:0/projects/test?token=secret").unwrap(); + let client = Client::new( + "http://user:password@127.0.0.1:0/projects/test?token=secret", + StoreAuth::Anonymous, + ) + .unwrap(); let error = client.fetch(b"k", b"s").await.unwrap_err(); assert!(matches!(error, Error::Network(_)), "{error:?}"); - let messages = - std::iter::successors(Some(&error as &dyn std::error::Error), |err| err.source()) - .map(|err| vt_str::format!("{err}")); - for text in messages.chain([vt_str::format!("{error:?}")]) { + for text in error_texts(&error) { assert!(!text.contains("password") && !text.contains("secret"), "{text}"); } } + + #[tokio::test] + async fn only_stores_send_an_oidc_token() { + let dir = tempfile::tempdir().unwrap(); + let blob = AbsolutePathBuf::new(dir.path().join("archive.tar.zst")).unwrap(); + std::fs::write(blob.as_path(), b"archive bytes").unwrap(); + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + let client = oidc_client_for(&listener); + let token = jwt(unix_now() + 3600); + let server = serve_each( + listener, + vec![ + http_response("HTTP/1.1 404 Not Found", b""), + stored(), + token_response(&token), + stored(), + ], + ); + + assert_eq!(client.fetch(b"k", b"s").await.unwrap(), None); + read_to_end(client.download("7").await.unwrap()).await.unwrap(); + client.store(b"k", b"s", b"v", Some(&blob)).await.unwrap(); + + let [fetch_request, download_request, token_request, store_request] = + server.join().unwrap().try_into().unwrap(); + assert!(!contains(&fetch_request, b"authorization:")); + assert!(!contains(&download_request, b"authorization:")); + let token_request_line = vt_str::format!( + "GET /token?api-version=2.0&audience=http%3A%2F%2F127.0.0.1%3A{port}%2Fprojects%2Ftest HTTP/1.1\r\n" + ); + assert!(token_request.starts_with(token_request_line.as_bytes())); + assert!(contains(&token_request, &authorization_line(REQUEST_TOKEN))); + assert!(is_store_request(&store_request)); + assert!(contains(&store_request, &authorization_line(&token))); + assert!(!contains(&store_request, REQUEST_TOKEN.as_bytes())); + assert!(contains(&store_request, b"archive bytes")); + + let store_auth = StoreAuth::GithubOidc(GithubOidc::new("https://token.example/", "secret")); + for text in [vt_str::format!("{client:?}"), vt_str::format!("{store_auth:?}")] { + for secret in [token.as_str(), REQUEST_TOKEN, "secret", "token.example", "127.0.0.1"] { + assert!(!text.contains(secret), "{text}"); + } + } + } + + #[tokio::test] + async fn oidc_token_is_reused_until_it_expires_soon() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let client = oidc_client_for(&listener); + // The first token expires within the refresh margin. + let expiring = jwt(unix_now() + 30); + let fresh = jwt(unix_now() + 3600); + let server = serve_each( + listener, + vec![token_response(&expiring), stored(), token_response(&fresh), stored(), stored()], + ); + + for _ in 0..3 { + client.store(b"k", b"s", b"v", None).await.unwrap(); + } + + let requests = server.join().unwrap(); + assert!(is_token_request(&requests[0])); + assert!(contains(&requests[1], &authorization_line(&expiring))); + assert!(is_token_request(&requests[2])); + assert!(contains(&requests[3], &authorization_line(&fresh))); + assert!(contains(&requests[4], &authorization_line(&fresh))); + } + + #[tokio::test] + async fn concurrent_stores_share_one_token_request() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let client = oidc_client_for(&listener); + let token = jwt(unix_now() + 3600); + // A second token request would get a store response, which isn't a + // token. + let server = serve_each(listener, vec![token_response(&token), stored(), stored()]); + + let (first, second) = tokio::join!( + client.store(b"k", b"s", b"v", None), + client.store(b"k", b"s", b"v", None), + ); + first.unwrap(); + second.unwrap(); + + let requests = server.join().unwrap(); + assert_eq!(requests.iter().filter(|request| is_token_request(request)).count(), 1); + for request in requests.iter().filter(|request| is_store_request(request)) { + assert!(contains(request, &authorization_line(&token))); + } + } + + #[tokio::test] + async fn failed_token_request_sends_no_store() { + let token_in_a_string = vt_str::format!(r#""{}""#, jwt(unix_now() + 3600)); + let payload = |json: &str| URL_SAFE_NO_PAD.encode(json); + let without_exp = vt_str::format!(r#"{{"value":"h.{}.s"}}"#, payload(r#"{"iat":1}"#)); + // The JSON escapes decode to a line break, which can't be in a header. + let with_line_break = + vt_str::format!(r#"{{"value":"h.{}.s\r\nx: y"}}"#, payload(r#"{"exp":4102444800}"#)); + for (response, expected) in [ + ( + http_response("HTTP/1.1 500 Internal Server Error", b"secret body"), + "HTTP status 500", + ), + ( + b"HTTP/1.1 302 Found\r\nlocation: /projects/test/store\r\ncontent-length: 0\r\nconnection: close\r\n\r\n" + .to_vec(), + "HTTP status 302", + ), + (http_response("HTTP/1.1 200 OK", b"secret body"), "malformed response"), + (http_response("HTTP/1.1 200 OK", token_in_a_string.as_bytes()), "malformed response"), + (http_response("HTTP/1.1 200 OK", br#"{"token":"secret"}"#), "malformed response"), + (http_response("HTTP/1.1 200 OK", br#"{"value":"secret"}"#), "malformed response"), + (http_response("HTTP/1.1 200 OK", without_exp.as_bytes()), "malformed response"), + (http_response("HTTP/1.1 200 OK", with_line_break.as_bytes()), "malformed response"), + ] { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let client = oidc_client_for(&listener); + // Only the token request is served. A store would fail to connect. + let server = serve_each(listener, vec![response]); + + let error = client.store(b"k", b"s", b"v", None).await.unwrap_err(); + assert!(matches!(error, Error::OidcToken(_)), "{error:?}"); + assert!(error.is_unauthorized()); + assert_eq!(error.to_string(), "failed to get a GitHub Actions OIDC token"); + assert_eq!(std::error::Error::source(&error).unwrap().to_string(), expected); + for text in error_texts(&error) { + // `eyJ` starts each part of the token, encoding `{"`. + for secret in ["secret", "eyJ", "/token", "127.0.0.1", REQUEST_TOKEN] { + assert!(!text.contains(secret), "{text}"); + } + } + assert!(is_token_request(&server.join().unwrap()[0])); + + // The failure is kept, so no other token request is made. + let again = client.store(b"k", b"s", b"v", None).await.unwrap_err(); + assert_eq!(error_texts(&again), error_texts(&error)); + } + } + + #[tokio::test] + async fn invalid_token_request_settings_send_no_request() { + for oidc in [ + GithubOidc::new("not a url", REQUEST_TOKEN), + GithubOidc::new("http://127.0.0.1:0/token", "secret\r\nx-injected: 1"), + ] { + // Nothing can listen on port 0, so a request would be a network + // error. + let client = + Client::new("http://127.0.0.1:0/projects/test", StoreAuth::GithubOidc(oidc)) + .unwrap(); + let error = client.store(b"k", b"s", b"v", None).await.unwrap_err(); + assert_eq!( + error_texts(&error)[..2], + ["failed to get a GitHub Actions OIDC token", "invalid token request"] + ); + } + } + + #[tokio::test] + async fn only_a_401_from_github_actions_without_oidc_has_a_hint() { + for (store_auth, status_line, expected) in [ + ( + StoreAuth::GithubActionsWithoutOidc, + "HTTP/1.1 401 Unauthorized", + ["HTTP status 401", "grant `id-token: write` to this job", "denied"].as_slice(), + ), + ( + StoreAuth::GithubActionsWithoutOidc, + "HTTP/1.1 403 Forbidden", + &["HTTP status 403", "denied"], + ), + (StoreAuth::Anonymous, "HTTP/1.1 401 Unauthorized", &["HTTP status 401", "denied"]), + ] { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let client = client_with_auth(&listener, store_auth); + let server = std::thread::spawn(move || serve_once(&listener, status_line, b"denied")); + + let error = client.store(b"k", b"s", b"v", None).await.unwrap_err(); + assert!(error.is_unauthorized()); + assert_eq!(error_texts(&error)[..expected.len()], *expected); + assert!(!contains(&server.join().unwrap(), b"authorization:")); + } + } }