diff --git a/CHANGELOG.md b/CHANGELOG.md index 61217dff1..34d209565 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,7 +4,7 @@ - **Changed** The run summary now says a task that wrote a file it also read was `not cached because it modified its inputs`, and the statistics in `vp run --verbose` and `vp run --last-details` use the singular for a count of one, e.g. `1 task • 1 cache miss` ([#783](https://github.com/voidzero-dev/vite-task/pull/783)). - **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. Uploads run in the background, so tasks that depend on an uploading task don't wait for it; once all tasks are done, `vp run` waits for the uploads still running and says so. A failed upload doesn't fail the task; the run summary shows a warning instead. Ctrl-C, or a failing task, stops remote cache lookups right away, and a task still being looked up doesn't start. Ctrl-C also cancels the uploads, but a failing task doesn't. Tasks can opt out with `cache: { remote: false }`. Requests to an HTTPS endpoint use HTTP/2 if the server supports it. 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), [#770](https://github.com/voidzero-dev/vite-task/pull/770), [#771](https://github.com/voidzero-dev/vite-task/pull/771), [#772](https://github.com/voidzero-dev/vite-task/pull/772), [#786](https://github.com/voidzero-dev/vite-task/pull/786), [#787](https://github.com/voidzero-dev/vite-task/pull/787)). +- **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. Uploads run in the background, so tasks that depend on an uploading task don't wait for it; once all tasks are done, `vp run` waits for the uploads still running and says so. A failed upload doesn't fail the task; the run summary shows a warning instead. Ctrl-C, or a failing task, stops remote cache lookups right away, and a task still being looked up doesn't start. Ctrl-C also cancels the uploads, but a failing task doesn't. Tasks can opt out with `cache: { remote: false }`. Requests to an HTTPS endpoint use HTTP/2 if the server supports it. Requests use the proxy environment variables or, on macOS and Windows, the system proxy settings. In a GitHub Actions job with `permissions: id-token: write`, uploads authenticate with a GitHub Actions OIDC token whose audience is the endpoint without a trailing slash. If no token can be obtained, each upload fails with a warning instead ([#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), [#770](https://github.com/voidzero-dev/vite-task/pull/770), [#771](https://github.com/voidzero-dev/vite-task/pull/771), [#772](https://github.com/voidzero-dev/vite-task/pull/772), [#786](https://github.com/voidzero-dev/vite-task/pull/786), [#787](https://github.com/voidzero-dev/vite-task/pull/787), [#798](https://github.com/voidzero-dev/vite-task/pull/798)). - **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 d0d3f76bd..ea866f6dc 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -5272,12 +5272,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/remote.rs b/crates/vt/src/session/cache/remote.rs index b5df5fffe..f45f874d5 100644 --- a/crates/vt/src/session/cache/remote.rs +++ b/crates/vt/src/session/cache/remote.rs @@ -30,7 +30,7 @@ use vt_plan::{ }; use vt_remote_cache::{ Client, Download, Fetched, - auth::{Anonymous, Auth}, + auth::{Anonymous, Auth, GithubOidc}, }; use vt_str::Str; use wincode::{ @@ -246,6 +246,11 @@ impl RemoteClients { fn build_auth(auth: &RemoteCacheAuth) -> Arc { match auth { RemoteCacheAuth::Anonymous => Arc::new(Anonymous), + RemoteCacheAuth::GithubOidc(github_oidc) => Arc::new(GithubOidc::new( + &github_oidc.request_url, + github_oidc.request_token.expose(), + &github_oidc.audience, + )), } } @@ -426,7 +431,7 @@ mod tests { use vt_path::{AbsolutePathBuf, RelativePathBuf}; use vt_plan::{ cache_metadata::{EnvValueHash, SpawnFingerprint}, - remote_cache::RemoteCacheAccess, + remote_cache::{GithubOidcAuth, RemoteCacheAccess, Secret}, }; use super::*; @@ -765,12 +770,22 @@ mod tests { serve_stalled(b"HTTP/1.1 500 Internal Server Error\r\ncontent-length: 0\r\n\r\n"); // Nothing can listen on port 0. let unreachable = anonymous("http://127.0.0.1:0/projects/test"); + let github_oidc_unavailable = ResolvedRemoteCacheConfig { + auth: RemoteCacheAuth::GithubOidc(GithubOidcAuth { + request_url: Arc::from("http://127.0.0.1:0/token"), + request_token: Secret::new(Arc::from("request-token")), + audience: Arc::from("http://127.0.0.1:0/projects/test"), + }), + ..unreachable.clone() + }; let clients = RemoteClients::default(); let uploads = RemoteUploads::default(); - for (remote_config, message) in - [(error_status, "HTTP status 500"), (unreachable, "network error")] - { + for (remote_config, message) in [ + (error_status, "HTTP status 500"), + (unreachable, "network error"), + (github_oidc_unavailable, "failed to authenticate"), + ] { let error = Arc::new(OnceLock::new()); uploads.spawn(prepare_upload(&clients, &remote_config).unwrap(), Arc::clone(&error)); uploads.wait(&CancellationToken::new()).await; 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 b57793feb..0647a836e 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 @@ -443,6 +443,42 @@ steps = [ ], comment = "The details include the underlying error." }, ] +[[e2e]] +name = "github_oidc_unavailable" +steps = [ + { argv = [ + "vt", + "run", + "build", + ], envs = [ + [ + "VP_REMOTE_CACHE", + "read-write", + ], + [ + "VP_REMOTE_CACHE_URL", + "http://127.0.0.1:0/projects/test", + ], + [ + "ACTIONS_ID_TOKEN_REQUEST_URL", + "http://127.0.0.1:0/token?api-version=2.0", + ], + [ + "ACTIONS_ID_TOKEN_REQUEST_TOKEN", + "request-token", + ], + [ + "VP_RUN_INTERNAL_HIDE_PENDING_UPLOADS", + "1", + ], + ], comment = "The job can request GitHub Actions OIDC tokens, so the upload needs one first. Nothing can listen on port 0, so the token request fails, and the upload fails without being sent. The task succeeds." }, + { argv = [ + "vt", + "run", + "--last-details", + ], comment = "The details include why the token request failed." }, +] + [[e2e]] name = "read_invalid_endpoint" steps = [ diff --git a/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_unavailable.md b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_unavailable.md new file mode 100644 index 000000000..e83ba35cd --- /dev/null +++ b/crates/vt_bin/tests/e2e_snapshots/fixtures/remote_cache/snapshots/github_oidc_unavailable.md @@ -0,0 +1,43 @@ +# github_oidc_unavailable + +## `VP_REMOTE_CACHE=read-write VP_REMOTE_CACHE_URL=http://127.0.0.1:0/projects/test ACTIONS_ID_TOKEN_REQUEST_URL=http://127.0.0.1:0/token?api-version=2.0 ACTIONS_ID_TOKEN_REQUEST_TOKEN=request-token VP_RUN_INTERNAL_HIDE_PENDING_UPLOADS=1 vt run build` + +The job can request GitHub Actions OIDC tokens, so the upload needs one first. Nothing can listen on port 0, so the token request fails, and the upload fails without being sent. The task succeeds. + +``` +$ vtt write-file dist/output.txt built ○ cache miss: remote cache fetch failed, executing + +--- +vt run: remote-cache#build not uploaded to the remote cache: failed to authenticate. (Run `vt run --last-details` for full details) +``` + +## `vt run --last-details` + +The details include why the token request failed. + +``` + +━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ + Vite+ Task Runner • Execution Summary +━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ + +Statistics: 1 task • 0 cache hits • 1 cache miss +Performance: 0% cache hit rate + +Task Details: +──────────────────────────────────────────────── + [1] remote-cache#build: $ vtt write-file dist/output.txt built ✓ + → Cache miss: remote cache fetch failed + ↳ network error + ↳ error sending request + ↳ client error (Connect) + ↳ tcp connect error + ↳ + ⚠ Not uploaded to the remote cache: failed to authenticate + ↳ GitHub Actions OIDC token request failed + ↳ error sending request + ↳ client error (Connect) + ↳ tcp connect error + ↳ +━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ +``` diff --git a/crates/vt_graph/run-config.ts b/crates/vt_graph/run-config.ts index b1d68ee16..d162f3bf8 100644 --- a/crates/vt_graph/run-config.ts +++ b/crates/vt_graph/run-config.ts @@ -29,6 +29,10 @@ export type InputBase = "package" | "workspace"; export type RemoteCacheConfig = { /** * HTTP or HTTPS namespace endpoint. Overridden by `VP_REMOTE_CACHE_URL`. + * + * In a GitHub Actions job with `permissions: id-token: write`, uploads + * authenticate with a GitHub Actions OIDC token whose audience is the + * endpoint without a trailing slash. */ url: string, }; diff --git a/crates/vt_graph/src/config/user.rs b/crates/vt_graph/src/config/user.rs index 0564a97eb..a0ef5888a 100644 --- a/crates/vt_graph/src/config/user.rs +++ b/crates/vt_graph/src/config/user.rs @@ -412,6 +412,10 @@ impl ResolvedGlobalCacheConfig { #[serde(deny_unknown_fields, rename_all = "camelCase")] pub struct UserRemoteCacheConfig { /// HTTP or HTTPS namespace endpoint. Overridden by `VP_REMOTE_CACHE_URL`. + /// + /// In a GitHub Actions job with `permissions: id-token: write`, uploads + /// authenticate with a GitHub Actions OIDC token whose audience is the + /// endpoint without a trailing slash. pub url: Arc, } diff --git a/crates/vt_plan/src/error.rs b/crates/vt_plan/src/error.rs index 39c5f87fe..6391fb0b8 100644 --- a/crates/vt_plan/src/error.rs +++ b/crates/vt_plan/src/error.rs @@ -184,6 +184,10 @@ pub enum Error { #[error("Remote caching requires cache.remote.url or VP_REMOTE_CACHE_URL")] MissingRemoteCacheEndpoint, + /// The value isn't shown, since it can be a credential. + #[error("Invalid value for {0}: not valid UTF-8")] + NonUtf8RemoteCacheAuthEnv(&'static str), + /// A cycle was detected in the task dependency graph during planning. /// /// This is caught by `AcyclicGraph::try_from_graph`, which validates that the diff --git a/crates/vt_plan/src/remote_cache.rs b/crates/vt_plan/src/remote_cache.rs index 8dcff2da6..26c9df37f 100644 --- a/crates/vt_plan/src/remote_cache.rs +++ b/crates/vt_plan/src/remote_cache.rs @@ -1,16 +1,20 @@ //! Remote cache settings: the mode requested with `--remote-cache` or //! `VP_REMOTE_CACHE`, and the access and auth resolved for each `vp run` level. -use std::{ffi::OsStr, sync::Arc}; +use std::{ffi::OsStr, fmt, sync::Arc}; use rustc_hash::FxHashMap; -use serde::Serialize; +use serde::{Serialize, Serializer}; use vt_casefold::EnvName; use crate::Error; pub(crate) const MODE_ENV: &str = "VP_REMOTE_CACHE"; const URL_ENV: &str = "VP_REMOTE_CACHE_URL"; +/// Set in GitHub Actions jobs that can request OIDC tokens, which need +/// `permissions: id-token: write`. +const GITHUB_OIDC_REQUEST_URL_ENV: &str = "ACTIONS_ID_TOKEN_REQUEST_URL"; +const GITHUB_OIDC_REQUEST_TOKEN_ENV: &str = "ACTIONS_ID_TOKEN_REQUEST_TOKEN"; /// Remote cache mode requested with `--remote-cache` or `VP_REMOTE_CACHE`. /// `resolve` combines it with the endpoint into a [`ResolvedRemoteCacheConfig`]. @@ -33,7 +37,8 @@ impl RemoteCacheMode { } /// Remote cache access, endpoint, and auth resolved for a `vp run` level from -/// the requested mode, `VP_REMOTE_CACHE_URL`, and `cache.remote.url`. +/// the requested mode, `VP_REMOTE_CACHE_URL`, `cache.remote.url`, and the envs +/// that the auth uses. #[derive(Debug, Clone, Serialize)] pub struct ResolvedRemoteCacheConfig { pub access: RemoteCacheAccess, @@ -49,6 +54,49 @@ pub struct ResolvedRemoteCacheConfig { pub enum RemoteCacheAuth { /// Requests carry no credentials. Anonymous, + /// Stores carry a GitHub Actions OIDC token. + GithubOidc(GithubOidcAuth), +} + +/// How stores get a GitHub Actions OIDC token: it's requested from +/// `request_url` with `request_token`, for `audience`. +#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize)] +pub struct GithubOidcAuth { + pub request_url: Arc, + pub request_token: Secret, + pub audience: Arc, +} + +/// A credential. Debug output and serialized plans show a placeholder +/// instead of its value. +#[derive(Clone, PartialEq, Eq, Hash)] +pub struct Secret(Arc); + +impl Secret { + const REDACTED: &str = ""; + + #[must_use] + pub const fn new(value: Arc) -> Self { + Self(value) + } + + /// The value, for sending it where it's needed. + #[must_use] + pub fn expose(&self) -> &str { + &self.0 + } +} + +impl fmt::Debug for Secret { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(Self::REDACTED) + } +} + +impl Serialize for Secret { + fn serialize(&self, serializer: S) -> Result { + serializer.serialize_str(Self::REDACTED) + } } /// Remote cache access after resolution. `off` resolves to no remote cache. @@ -98,25 +146,57 @@ pub(crate) fn resolve( return Err(Error::MissingRemoteCacheEndpoint); } }; - Ok(Some(ResolvedRemoteCacheConfig { access, url, auth: RemoteCacheAuth::Anonymous })) + let auth = resolve_auth(&url, envs)?; + Ok(Some(ResolvedRemoteCacheConfig { access, url, auth })) +} + +/// Resolves how requests to `url` authenticate from the envs visible at the +/// `vp run` level. In a GitHub Actions job that can request OIDC tokens, +/// stores carry a token whose audience is `url` without a trailing slash. +/// Otherwise requests are anonymous. +fn resolve_auth( + url: &str, + envs: &FxHashMap>, Arc>, +) -> Result { + let (Some(request_url), Some(request_token)) = ( + env_value(envs, GITHUB_OIDC_REQUEST_URL_ENV), + env_value(envs, GITHUB_OIDC_REQUEST_TOKEN_ENV), + ) else { + return Ok(RemoteCacheAuth::Anonymous); + }; + let utf8 = |name: &'static str, value: &OsStr| { + value.to_str().map(Arc::from).ok_or(Error::NonUtf8RemoteCacheAuthEnv(name)) + }; + Ok(RemoteCacheAuth::GithubOidc(GithubOidcAuth { + request_url: utf8(GITHUB_OIDC_REQUEST_URL_ENV, request_url)?, + request_token: Secret::new(utf8(GITHUB_OIDC_REQUEST_TOKEN_ENV, request_token)?), + audience: Arc::from(url.trim_end_matches('/')), + })) } #[cfg(test)] mod tests { + use std::ffi::OsString; + use super::*; + fn envs<'a>( + pairs: impl IntoIterator, + ) -> FxHashMap>, Arc> { + pairs + .into_iter() + .map(|(name, value)| { + (EnvName::new(Arc::::from(OsStr::new(name))), Arc::::from(value)) + }) + .collect() + } + #[test] fn reads_envs_by_platform_name_rules() { - let envs = - [("vp_remote_cache", "read-write"), ("vp_remote_cache_url", "https://cache.example")] - .into_iter() - .map(|(name, value)| { - ( - EnvName::new(Arc::::from(OsStr::new(name))), - Arc::::from(OsStr::new(value)), - ) - }) - .collect(); + let envs = envs([ + ("vp_remote_cache", OsStr::new("read-write")), + ("vp_remote_cache_url", OsStr::new("https://cache.example")), + ]); let resolved = resolve(None, &envs).unwrap(); if cfg!(windows) { let resolved = resolved.unwrap(); @@ -126,4 +206,50 @@ mod tests { assert!(resolved.is_none()); } } + #[test] + fn github_oidc_token_is_redacted() { + let envs = envs([ + ("VP_REMOTE_CACHE_URL", OsStr::new("https://cache.example/projects/test/")), + ("ACTIONS_ID_TOKEN_REQUEST_URL", OsStr::new("https://token.example/?api-version=2.0")), + ("ACTIONS_ID_TOKEN_REQUEST_TOKEN", OsStr::new("request-token")), + ]); + let auth = resolve(None, &envs).unwrap().unwrap().auth; + let RemoteCacheAuth::GithubOidc(github_oidc) = &auth else { + panic!("expected GitHub OIDC, got {auth:?}"); + }; + assert_eq!(&*github_oidc.request_url, "https://token.example/?api-version=2.0"); + assert_eq!(github_oidc.request_token.expose(), "request-token"); + assert_eq!(&*github_oidc.audience, "https://cache.example/projects/test"); + + let debug = vt_str::format!("{auth:?}"); + let serialized = serde_json::to_string(&auth).unwrap(); + for output in [debug.as_str(), &serialized] { + assert!(!output.contains("request-token") && output.contains(""), "{output}"); + } + } + + #[test] + fn non_utf8_auth_env_fails_without_showing_its_value() { + #[cfg(unix)] + let token = { + use std::os::unix::ffi::OsStringExt as _; + OsString::from_vec(b"secret\xff".to_vec()) + }; + #[cfg(windows)] + let token = { + use std::os::windows::ffi::OsStringExt as _; + // "secret" followed by an unpaired surrogate. + OsString::from_wide(&[0x73, 0x65, 0x63, 0x72, 0x65, 0x74, 0xD800]) + }; + let envs = envs([ + ("VP_REMOTE_CACHE_URL", OsStr::new("https://cache.example")), + ("ACTIONS_ID_TOKEN_REQUEST_URL", OsStr::new("https://token.example")), + ("ACTIONS_ID_TOKEN_REQUEST_TOKEN", token.as_os_str()), + ]); + let error = resolve(None, &envs).unwrap_err(); + assert_eq!( + error.to_string(), + "Invalid value for ACTIONS_ID_TOKEN_REQUEST_TOKEN: not valid UTF-8" + ); + } } diff --git a/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots.toml b/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots.toml index bf023f6b1..9d2f16b56 100644 --- a/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots.toml +++ b/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots.toml @@ -70,3 +70,17 @@ args = ["run", "siblings"] [[plan]] name = "sibling_endpoint_isolation" args = ["run", "endpoints"] + +# In a GitHub Actions job that can request OIDC tokens, stores authenticate +# with a token whose audience is the endpoint without a trailing slash. The +# request token is redacted. +[[plan]] +name = "github_oidc_auth" +args = ["run", "build"] +env = { "VP_REMOTE_CACHE_URL" = "https://cache.example/projects/test/", "ACTIONS_ID_TOKEN_REQUEST_URL" = "https://token.example/idtoken?api-version=2.0", "ACTIONS_ID_TOKEN_REQUEST_TOKEN" = "request-token" } + +# Without the request token, requests are anonymous. +[[plan]] +name = "github_oidc_auth_needs_request_token" +args = ["run", "build"] +env = { "ACTIONS_ID_TOKEN_REQUEST_URL" = "https://token.example/idtoken?api-version=2.0" } diff --git a/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots/query_github_oidc_auth.jsonc b/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots/query_github_oidc_auth.jsonc new file mode 100644 index 000000000..1f54c163e --- /dev/null +++ b/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots/query_github_oidc_auth.jsonc @@ -0,0 +1,48 @@ +// run build +{ + "graph": [ + { + "key": [ + "/", + "build" + ], + "node": { + "items": [ + { + "kind": { + "Leaf": { + "Spawn": { + "cache_metadata": { + "spawn_fingerprint": { + "env_fingerprints": { + "fingerprinted_envs": {}, + "untracked_env_config": [ + "" + ] + } + }, + "remote_cache": { + "access": "read", + "url": "https://cache.example/projects/test/", + "auth": { + "kind": "github-oidc", + "request_url": "https://token.example/idtoken?api-version=2.0", + "request_token": "", + "audience": "https://cache.example/projects/test" + } + } + }, + "spawn_command": { + "spawn_envs": { + "VP_REMOTE_CACHE_URL": "https://cache.example/projects/test/" + } + } + } + } + } + } + ] + } + } + ] +} diff --git a/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots/query_github_oidc_auth_needs_request_token.jsonc b/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots/query_github_oidc_auth_needs_request_token.jsonc new file mode 100644 index 000000000..9d5bbc8ff --- /dev/null +++ b/crates/vt_plan/tests/plan_snapshots/fixtures/remote_cache_config/snapshots/query_github_oidc_auth_needs_request_token.jsonc @@ -0,0 +1,43 @@ +// run build +{ + "graph": [ + { + "key": [ + "/", + "build" + ], + "node": { + "items": [ + { + "kind": { + "Leaf": { + "Spawn": { + "cache_metadata": { + "spawn_fingerprint": { + "env_fingerprints": { + "fingerprinted_envs": {}, + "untracked_env_config": [ + "" + ] + } + }, + "remote_cache": { + "access": "read", + "url": "https://cache.example/projects/test", + "auth": { + "kind": "anonymous" + } + } + }, + "spawn_command": { + "spawn_envs": {} + } + } + } + } + } + ] + } + } + ] +} diff --git a/crates/vt_remote_cache/Cargo.toml b/crates/vt_remote_cache/Cargo.toml index e8bcd52e3..13d7e8b26 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 = [ @@ -20,8 +21,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 60e4d6e3a..86ddb1816 100644 --- a/crates/vt_remote_cache/README.md +++ b/crates/vt_remote_cache/README.md @@ -6,6 +6,8 @@ Client for the [remote cache server API](https://github.com/voidzero-dev/vite-ta `Client::new` also takes an `Auth`, which supplies the headers that authenticate each request. It's asked before every request with the request's operation (fetch, download, or store), so it can authenticate some operations and not others. It can use the client's HTTP client to get its credentials, such as a token. If it fails, the request isn't sent, and the operation fails with `Error::Auth`. `Anonymous` adds no headers. +`GithubOidc` authenticates stores with a GitHub Actions OIDC token, sent as `Authorization: Bearer `. Fetches and downloads carry no credentials. It takes the job's token request URL and request token, the values of `ACTIONS_ID_TOKEN_REQUEST_URL` and `ACTIONS_ID_TOKEN_REQUEST_TOKEN`, and the token's audience, but reads no envs itself. The first store requests a token from the request URL with the audience added to its query, and later stores reuse it until two minutes before the expiry in its `exp` claim, which isn't verified. The margin covers the upload, since a server can check the token only once the whole store has arrived, as Cloudflare Workers do. A token without an `exp` claim is a malformed response. Concurrent stores wait for the same token request. Once a request fails, every later store fails the same way without another request. If the request URL isn't an HTTP or HTTPS URL, or the request token can't be sent in a header, every store fails. Neither token appears in its debug output or errors. + `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. `Client::download` gets a blob by its ID. After a 200 response, it returns a `Download`, from which the caller reads the blob's chunks as they arrive. diff --git a/crates/vt_remote_cache/src/auth/github_oidc.rs b/crates/vt_remote_cache/src/auth/github_oidc.rs new file mode 100644 index 000000000..b705ee1b1 --- /dev/null +++ b/crates/vt_remote_cache/src/auth/github_oidc.rs @@ -0,0 +1,491 @@ +//! Authenticates stores with GitHub Actions OIDC tokens. + +use std::{ + fmt, + sync::Arc, + time::{Duration, SystemTime}, +}; + +use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD}; +use reqwest::{ + StatusCode, + header::{ACCEPT, AUTHORIZATION, HeaderMap, HeaderValue}, +}; +use serde::Deserialize; +use tokio::sync::Mutex; +use url::{ParseError, Url}; +use vt_str::Str; + +use super::{Auth, AuthError, AuthHeaders, Operation}; + +/// Time allowed for a token request, from sending it until its whole response +/// arrives. +const TOKEN_REQUEST_TIMEOUT: Duration = Duration::from_secs(10); + +/// A token that expires within this time isn't reused. A server can check the +/// token only once the whole store has arrived, as Cloudflare Workers do, and +/// check it again when it publishes the store, so the token has to outlast the +/// upload and the server's work on it. +const REFRESH_MARGIN: Duration = Duration::from_secs(120); + +/// Authenticates stores with a GitHub Actions OIDC token, sent as +/// `Authorization: Bearer `. Fetches and downloads carry no +/// credentials. +/// +/// The token is requested when the first store needs it, and later stores +/// reuse it until shortly before it expires. Concurrent stores wait for the +/// same request. Once a request fails, every later store fails the same way, +/// without another request. +pub struct GithubOidc { + state: Mutex, +} + +/// What the next store does. +enum State { + /// It reuses `token` if it's still fresh, and otherwise requests a new + /// one with `request`. + Ready { request: TokenRequest, token: Option }, + /// It fails with this error, without a request. + Failed(AuthError), +} + +impl GithubOidc { + /// Authenticate stores with tokens for `audience`, requested from + /// `request_url` with `request_token`. In a GitHub Actions job, those are + /// the values of `ACTIONS_ID_TOKEN_REQUEST_URL` and + /// `ACTIONS_ID_TOKEN_REQUEST_TOKEN`. Nothing is requested yet. If + /// `request_url` isn't an HTTP or HTTPS URL, or `request_token` can't be + /// sent in a header, every store fails. + #[must_use] + pub fn new(request_url: &str, request_token: &str, audience: &str) -> Self { + let state = match TokenRequest::new(request_url, request_token, audience) { + Ok(request) => State::Ready { request, token: None }, + Err(err) => State::Failed(Arc::new(err)), + }; + Self { state: Mutex::new(state) } + } + + /// The `Authorization` header for a store. + async fn authorization(&self, http: &reqwest::Client) -> Result { + // The lock is held during the request, so concurrent stores wait for + // its token instead of making their own requests. + let mut state = self.state.lock().await; + let (request, token) = match &mut *state { + State::Ready { request, token } => (request, token), + State::Failed(err) => return Err(Arc::clone(err)), + }; + if let Some(token) = token.as_ref().filter(|token| SystemTime::now() < token.refresh_at) { + return Ok(token.authorization.clone()); + } + let authorization = match request.send(http).await { + Ok(fresh) => Ok(token.insert(fresh).authorization.clone()), + Err(err) => { + let err: AuthError = Arc::new(err); + *state = State::Failed(Arc::clone(&err)); + Err(err) + } + }; + drop(state); + authorization + } +} + +impl fmt::Debug for GithubOidc { + /// Leaves out the tokens. + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("GithubOidc").finish_non_exhaustive() + } +} + +impl Auth for GithubOidc { + fn headers<'a>(&'a self, operation: Operation, http: &'a reqwest::Client) -> AuthHeaders<'a> { + Box::pin(async move { + let mut headers = HeaderMap::new(); + if operation == Operation::Store { + headers.insert(AUTHORIZATION, self.authorization(http).await?); + } + Ok(headers) + }) + } +} + +/// Why no token could be obtained. The messages name only the kind of +/// failure, and none of them contain a token. +#[derive(Debug, thiserror::Error)] +enum TokenError { + /// The source is the parse error if the request URL isn't a URL at all. + #[error("invalid GitHub Actions OIDC token request URL")] + InvalidRequestUrl(#[source] Option), + #[error("invalid GitHub Actions OIDC request token")] + InvalidRequestToken, + /// No complete response arrived. The source leaves out the request URL. + #[error("GitHub Actions OIDC token request failed")] + Network(#[source] reqwest::Error), + #[error("GitHub Actions OIDC token request failed with HTTP status {}", .0.as_u16())] + Status(StatusCode), + /// The body isn't JSON with a token in `value`, the token has no `exp` + /// claim, or it can't be sent in a header. + #[error("malformed GitHub Actions OIDC token response")] + MalformedResponse, +} + +/// How to request a token. +struct TokenRequest { + /// The request URL, with the audience added to its query. + url: Url, + /// `Bearer `. + authorization: HeaderValue, +} + +/// A token from a successful request. +struct Token { + /// `Bearer `. + authorization: HeaderValue, + /// When stores stop reusing the token: shortly before it expires. + refresh_at: SystemTime, +} + +/// The body of a successful token response. +#[derive(Deserialize)] +struct TokenResponse { + value: Str, +} + +impl TokenRequest { + fn new(request_url: &str, request_token: &str, audience: &str) -> Result { + let mut url = + Url::parse(request_url).map_err(|err| TokenError::InvalidRequestUrl(Some(err)))?; + if !matches!(url.scheme(), "http" | "https") { + return Err(TokenError::InvalidRequestUrl(None)); + } + url.query_pairs_mut().append_pair("audience", audience); + let authorization = bearer(request_token).ok_or(TokenError::InvalidRequestToken)?; + Ok(Self { url, authorization }) + } + + async fn send(&self, http: &reqwest::Client) -> Result { + let network_error = |err: reqwest::Error| TokenError::Network(err.without_url()); + let response = http + .get(self.url.clone()) + .header(AUTHORIZATION, self.authorization.clone()) + .header(ACCEPT, "application/json") + .timeout(TOKEN_REQUEST_TIMEOUT) + .send() + .await + .map_err(network_error)?; + let status = response.status(); + if status != StatusCode::OK { + return Err(TokenError::Status(status)); + } + let body = response.bytes().await.map_err(network_error)?; + let TokenResponse { value } = + serde_json::from_slice(&body).map_err(|_| TokenError::MalformedResponse)?; + let refresh_at = expiry(&value) + .and_then(|expiry| expiry.checked_sub(REFRESH_MARGIN)) + .ok_or(TokenError::MalformedResponse)?; + let authorization = bearer(&value).ok_or(TokenError::MalformedResponse)?; + Ok(Token { authorization, refresh_at }) + } +} + +/// `Bearer `, marked sensitive so that debug output leaves it out, or +/// `None` if `token` can't be sent in a header. +fn bearer(token: &str) -> Option { + let mut value = HeaderValue::from_str(&vt_str::format!("Bearer {token}")).ok()?; + value.set_sensitive(true); + Some(value) +} + +/// When `token`, a JSON Web Token, expires according to its `exp` claim. The +/// claim isn't verified, since it only decides when to request a new token. +fn expiry(token: &str) -> Option { + #[derive(Deserialize)] + struct Claims { + exp: u64, + } + + let payload = token.split('.').nth(1)?; + let payload = URL_SAFE_NO_PAD.decode(payload).ok()?; + let Claims { exp } = serde_json::from_slice(&payload).ok()?; + SystemTime::UNIX_EPOCH.checked_add(Duration::from_secs(exp)) +} + +#[cfg(test)] +mod tests { + use std::{net::TcpListener, thread::JoinHandle}; + + use super::*; + use crate::{ + Client, Error, + test_server::{contains, no_request_waiting, serve_once}, + }; + + const AUDIENCE: &str = "https://cache.example/projects/test"; + + /// A JSON Web Token that expires `lifetime` from now. Only its payload is + /// read; `id` tells tokens apart. + fn jwt(lifetime: Duration, id: &str) -> Str { + let exp = (SystemTime::now() + lifetime).duration_since(SystemTime::UNIX_EPOCH).unwrap(); + let payload = + URL_SAFE_NO_PAD.encode(vt_str::format!("{{\"exp\":{}}}", exp.as_secs()).as_bytes()); + vt_str::format!("eyJhbGciOiJSUzI1NiJ9.{payload}.{id}") + } + + fn token_response(token: &str) -> Vec { + vt_str::format!("{{\"count\":1,\"value\":\"{token}\"}}").as_bytes().to_vec() + } + + /// A token endpoint and a cache endpoint on loopback, and a client for + /// the cache endpoint whose stores get tokens from the token endpoint. + struct Endpoints { + token: Arc, + cache: Arc, + auth: Arc, + client: Client, + } + + impl Endpoints { + fn new() -> Self { + let token = TcpListener::bind("127.0.0.1:0").unwrap(); + let cache = TcpListener::bind("127.0.0.1:0").unwrap(); + let request_url = + vt_str::format!("http://{}/token?api-version=2.0", token.local_addr().unwrap()); + let auth = Arc::new(GithubOidc::new(&request_url, "request-token", AUDIENCE)); + let endpoint = vt_str::format!("http://{}/projects/test", cache.local_addr().unwrap()); + let client = Client::new(&endpoint, Arc::clone(&auth) as Arc).unwrap(); + Self { token: Arc::new(token), cache: Arc::new(cache), auth, client } + } + + /// Serve token requests in turn, responding with `responses`. + fn serve_tokens( + &self, + responses: Vec<(&'static str, Vec)>, + ) -> JoinHandle>> { + serve(&self.token, responses) + } + + /// Serve `count` stores in turn, each responding with 200. + fn serve_stores(&self, count: usize) -> JoinHandle>> { + serve(&self.cache, vec![("HTTP/1.1 200 OK", Vec::new()); count]) + } + + async fn store(&self) -> Result<(), Error> { + self.client.store(b"k", b"s", b"v", None).await + } + } + + /// Serve a request for each of `responses` in turn, and return the + /// requests. + fn serve( + listener: &Arc, + responses: Vec<(&'static str, Vec)>, + ) -> JoinHandle>> { + let listener = Arc::clone(listener); + std::thread::spawn(move || { + responses.iter().map(|(status, body)| serve_once(&listener, status, body)).collect() + }) + } + + /// The error's message, followed by its sources' messages. + fn messages(error: &Error) -> 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 store_carries_a_token_for_the_audience() { + let endpoints = Endpoints::new(); + let token = jwt(Duration::from_secs(300), "a"); + let tokens = endpoints.serve_tokens(vec![("HTTP/1.1 200 OK", token_response(&token))]); + let stores = endpoints.serve_stores(1); + + endpoints.store().await.unwrap(); + + let token_request = &tokens.join().unwrap()[0]; + let target = "GET /token?api-version=2.0&audience=https%3A%2F%2Fcache.example%2Fprojects%2Ftest HTTP/1.1\r\n"; + assert!(token_request.starts_with(target.as_bytes())); + assert!(contains(token_request, b"authorization: Bearer request-token\r\n")); + assert!(contains(token_request, b"accept: application/json\r\n")); + let store_request = &stores.join().unwrap()[0]; + let authorization = vt_str::format!("authorization: Bearer {token}\r\n"); + assert!(contains(store_request, authorization.as_bytes())); + + let debug = vt_str::format!("{:?}", endpoints.auth); + assert!(!debug.contains("request-token") && !debug.contains(token.as_str()), "{debug}"); + } + + #[tokio::test] + async fn fetches_and_downloads_carry_no_token() { + let endpoints = Endpoints::new(); + let cache = Arc::clone(&endpoints.cache); + let server = std::thread::spawn(move || { + [ + serve_once(&cache, "HTTP/1.1 404 Not Found", b""), + serve_once(&cache, "HTTP/1.1 404 Not Found", b""), + ] + }); + + assert_eq!(endpoints.client.fetch(b"k", b"s").await.unwrap(), None); + endpoints.client.download("7").await.unwrap_err(); + + for request in server.join().unwrap() { + assert!(!contains(&request, b"authorization:")); + } + assert!(no_request_waiting(&endpoints.token)); + } + + #[tokio::test] + async fn stores_reuse_a_token_until_shortly_before_it_expires() { + let endpoints = Endpoints::new(); + let token = jwt(Duration::from_secs(300), "a"); + let tokens = endpoints.serve_tokens(vec![("HTTP/1.1 200 OK", token_response(&token))]); + let stores = endpoints.serve_stores(2); + + endpoints.store().await.unwrap(); + endpoints.store().await.unwrap(); + + tokens.join().unwrap(); + assert!(no_request_waiting(&endpoints.token)); + let authorization = vt_str::format!("authorization: Bearer {token}\r\n"); + for request in stores.join().unwrap() { + assert!(contains(&request, authorization.as_bytes())); + } + } + + #[tokio::test] + async fn concurrent_stores_wait_for_one_token_request() { + let endpoints = Endpoints::new(); + let token = jwt(Duration::from_secs(300), "a"); + let tokens = endpoints.serve_tokens(vec![("HTTP/1.1 200 OK", token_response(&token))]); + let stores = endpoints.serve_stores(2); + + let (first, second) = tokio::join!(endpoints.store(), endpoints.store()); + first.unwrap(); + second.unwrap(); + + tokens.join().unwrap(); + assert!(no_request_waiting(&endpoints.token)); + stores.join().unwrap(); + } + + #[tokio::test] + async fn token_that_expires_soon_is_not_reused() { + let endpoints = Endpoints::new(); + let first = jwt(Duration::from_secs(100), "a"); + let second = jwt(Duration::from_secs(300), "b"); + let tokens = endpoints.serve_tokens(vec![ + ("HTTP/1.1 200 OK", token_response(&first)), + ("HTTP/1.1 200 OK", token_response(&second)), + ]); + let stores = endpoints.serve_stores(2); + + endpoints.store().await.unwrap(); + endpoints.store().await.unwrap(); + + assert_eq!(tokens.join().unwrap().len(), 2); + let stores = stores.join().unwrap(); + for (request, token) in stores.iter().zip([&first, &second]) { + let authorization = vt_str::format!("authorization: Bearer {token}\r\n"); + assert!(contains(request, authorization.as_bytes())); + } + } + + #[tokio::test] + async fn failed_token_request_fails_later_stores_without_requests() { + let endpoints = Endpoints::new(); + let tokens = endpoints.serve_tokens(vec![("HTTP/1.1 403 Forbidden", b"denied".to_vec())]); + + let first = endpoints.store().await.unwrap_err(); + tokens.join().unwrap(); + let second = endpoints.store().await.unwrap_err(); + + for error in [first, second] { + assert!(matches!(error, Error::Auth(_)), "{error:?}"); + assert_eq!( + messages(&error), + [ + "failed to authenticate", + "GitHub Actions OIDC token request failed with HTTP status 403" + ] + ); + } + assert!(no_request_waiting(&endpoints.token)); + assert!(no_request_waiting(&endpoints.cache)); + } + + #[tokio::test] + async fn malformed_token_responses_fail_the_store() { + for body in [ + b"not json".to_vec(), + b"{}".to_vec(), + br#"{"value":1}"#.to_vec(), + br#"{"value":""}"#.to_vec(), + // Tokens without an `exp` claim. + br#"{"value":"opaque"}"#.to_vec(), + br#"{"value":"eyJhbGciOiJSUzI1NiJ9.e30.a"}"#.to_vec(), + // A token with a newline, escaped in JSON, can't be sent in a header. + token_response(&jwt(Duration::from_secs(300), "a\\nb")), + ] { + let endpoints = Endpoints::new(); + let tokens = endpoints.serve_tokens(vec![("HTTP/1.1 200 OK", body)]); + + let error = endpoints.store().await.unwrap_err(); + assert_eq!( + messages(&error), + ["failed to authenticate", "malformed GitHub Actions OIDC token response"] + ); + tokens.join().unwrap(); + assert!(no_request_waiting(&endpoints.cache)); + } + } + + #[tokio::test] + async fn unreachable_token_endpoint_is_a_network_error() { + let cache = TcpListener::bind("127.0.0.1:0").unwrap(); + // Nothing can listen on port 0. + let auth = GithubOidc::new("http://127.0.0.1:0/token", "request-token", AUDIENCE); + let endpoint = vt_str::format!("http://{}/projects/test", cache.local_addr().unwrap()); + let client = Client::new(&endpoint, Arc::new(auth)).unwrap(); + + let error = client.store(b"k", b"s", b"v", None).await.unwrap_err(); + let messages = messages(&error); + assert_eq!( + messages[..2], + ["failed to authenticate", "GitHub Actions OIDC token request failed"] + ); + for message in &messages { + assert!( + !message.contains("127.0.0.1:0/token") && !message.contains("request-token"), + "{message}" + ); + } + assert!(no_request_waiting(&cache)); + } + + #[tokio::test] + async fn invalid_settings_fail_stores_but_not_fetches() { + for (request_url, request_token, message) in [ + ("not a url", "request-token", "invalid GitHub Actions OIDC token request URL"), + ( + "ftp://token.example/", + "request-token", + "invalid GitHub Actions OIDC token request URL", + ), + ("http://token.example/", "a\nb", "invalid GitHub Actions OIDC request token"), + ] { + let cache = Arc::new(TcpListener::bind("127.0.0.1:0").unwrap()); + let auth = GithubOidc::new(request_url, request_token, AUDIENCE); + let endpoint = vt_str::format!("http://{}/projects/test", cache.local_addr().unwrap()); + let client = Client::new(&endpoint, Arc::new(auth)).unwrap(); + + let error = client.store(b"k", b"s", b"v", None).await.unwrap_err(); + assert_eq!(messages(&error)[..2], ["failed to authenticate", message]); + assert!(no_request_waiting(&cache)); + + let server = serve(&cache, vec![("HTTP/1.1 404 Not Found", Vec::new())]); + assert_eq!(client.fetch(b"k", b"s").await.unwrap(), None); + server.join().unwrap(); + } + } +} diff --git a/crates/vt_remote_cache/src/auth/mod.rs b/crates/vt_remote_cache/src/auth/mod.rs index c86ee222b..b4dbab6a4 100644 --- a/crates/vt_remote_cache/src/auth/mod.rs +++ b/crates/vt_remote_cache/src/auth/mod.rs @@ -5,6 +5,10 @@ use std::{error::Error as StdError, fmt, pin::Pin, sync::Arc}; use reqwest::header::HeaderMap; +mod github_oidc; + +pub use github_oidc::GithubOidc; + /// The operation a request performs. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Operation { diff --git a/crates/vt_remote_cache/src/lib.rs b/crates/vt_remote_cache/src/lib.rs index 90349ab74..759792bb0 100644 --- a/crates/vt_remote_cache/src/lib.rs +++ b/crates/vt_remote_cache/src/lib.rs @@ -2,6 +2,8 @@ //! bytes; the caller decides what they contain. pub mod auth; +#[cfg(test)] +mod test_server; use std::{sync::Arc, time::Duration}; @@ -314,16 +316,16 @@ async fn blob_part(path: &AbsolutePath) -> Result { #[cfg(test)] mod tests { - use std::{ - io::{Read as _, Write as _}, - net::TcpListener, - }; + use std::{io::Read as _, net::TcpListener}; use reqwest::header::{HeaderName, HeaderValue}; use vt_path::AbsolutePathBuf; use super::*; - use crate::auth::{Anonymous, AuthHeaders}; + use crate::{ + auth::{Anonymous, AuthHeaders}, + test_server::{contains, no_request_waiting, serve_once, serve_raw_once}, + }; fn store_url(endpoint: &str) -> Result { route_url(&parse_endpoint(endpoint)?, "store") @@ -425,48 +427,6 @@ mod tests { } } - fn contains(haystack: &[u8], needle: &[u8]) -> bool { - haystack.windows(needle.len()).any(|window| window == needle) - } - - /// 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()) - } - - /// Accept one HTTP request, write `response`, close the connection, and - /// return the raw request. - fn serve_raw_once(listener: &TcpListener, response: &[u8]) -> Vec { - let (mut stream, _) = listener.accept().unwrap(); - let mut request = Vec::new(); - let mut buf = [0; 4096]; - let header_end = loop { - let n = stream.read(&mut buf).unwrap(); - assert_ne!(n, 0, "connection closed before the request headers 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[..header_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() < header_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]); - } - stream.write_all(response).unwrap(); - request - } - fn client_for(listener: &TcpListener) -> Client { client_with_auth(listener, Arc::new(Anonymous)) } @@ -738,8 +698,7 @@ mod tests { assert_eq!(error.to_string(), "failed to authenticate"); assert_eq!(std::error::Error::source(&error).unwrap().to_string(), "no credentials"); } - listener.set_nonblocking(true).unwrap(); - assert_eq!(listener.accept().unwrap_err().kind(), std::io::ErrorKind::WouldBlock); + assert!(no_request_waiting(&listener)); } /// Accept one connection and return the TLS record it starts with, which diff --git a/crates/vt_remote_cache/src/test_server.rs b/crates/vt_remote_cache/src/test_server.rs new file mode 100644 index 000000000..faf77ce27 --- /dev/null +++ b/crates/vt_remote_cache/src/test_server.rs @@ -0,0 +1,61 @@ +//! A loopback HTTP server for tests, which serves one request at a time. + +use std::{ + io::{Read as _, Write as _}, + net::TcpListener, +}; + +pub fn contains(haystack: &[u8], needle: &[u8]) -> bool { + haystack.windows(needle.len()).any(|window| window == needle) +} + +/// Accept one HTTP request, respond with `status_line` and `body`, and +/// return the raw request. The response closes the connection, so the client +/// sends its next request on a new one. +pub fn serve_once(listener: &TcpListener, status_line: &str, body: &[u8]) -> Vec { + let headers = vt_str::format!( + "{status_line}\r\nconnection: close\r\ncontent-length: {}\r\n\r\n", + body.len() + ); + serve_raw_once(listener, &[headers.as_bytes(), body].concat()) +} + +/// Accept one HTTP request, write `response`, close the connection, and +/// return the raw request. +pub fn serve_raw_once(listener: &TcpListener, response: &[u8]) -> Vec { + let (mut stream, _) = listener.accept().unwrap(); + let mut request = Vec::new(); + let mut buf = [0; 4096]; + let header_end = loop { + let n = stream.read(&mut buf).unwrap(); + assert_ne!(n, 0, "connection closed before the request headers 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[..header_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() < header_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]); + } + stream.write_all(response).unwrap(); + request +} + +/// Whether no connection is waiting on `listener`, so no request was sent to +/// it since the last one served. +pub fn no_request_waiting(listener: &TcpListener) -> bool { + listener.set_nonblocking(true).unwrap(); + let waiting = listener.accept(); + listener.set_nonblocking(false).unwrap(); + waiting.is_err_and(|err| err.kind() == std::io::ErrorKind::WouldBlock) +}