diff --git a/CHANGELOG.md b/CHANGELOG.md index 97e7e7da0..3572a298c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -231,6 +231,11 @@ limits, and required install commands. project and rewired its lockfile, and `vendor --revert -g` unwound the project's vendoring, so its next frozen install was silently unpatched. Global installs have no project lockfile to vendor into (#498). +- Patch API requests (`scan`, `get`, `apply` and `vex` lookups, and blob and + diff downloads) no longer hang forever on a stalled proxy, load balancer or + half-open connection. A connect now fails after 10 s, and a connection that + sends nothing for 60 s fails as a network error. Downloads that keep + streaming are not cut off (#570). ### Maintenance diff --git a/crates/socket-patch-bench/src/fixtures/npm.rs b/crates/socket-patch-bench/src/fixtures/npm.rs index 1985f703a..af1ee68e4 100644 --- a/crates/socket-patch-bench/src/fixtures/npm.rs +++ b/crates/socket-patch-bench/src/fixtures/npm.rs @@ -623,7 +623,13 @@ pub fn build_yarn_berry(t: &mut Tree, size: Size) -> std::io::Result { t.write("project/node_modules/.yarn-state.yml", "# Warning: This file is automatically generated. Removing it is fine, but will\n# cause your node_modules installation to become invalidated.\n\n__metadata:\n version: 1\n nmMode: classic\n")?; g.install_hoisted(t, "project/")?; t.mkdir("home")?; - Ok(fixture(&g, g.patches(true), &["yarn.lock"], &[])) + // Hosted Berry pins both descriptor resolutions and their lock entries. + Ok(fixture( + &g, + g.patches(true), + &["package.json", "yarn.lock"], + &[], + )) } // ── bun ──────────────────────────────────────────────────────────────── diff --git a/crates/socket-patch-core/src/api/client.rs b/crates/socket-patch-core/src/api/client.rs index b1d1174e9..7d862c050 100644 --- a/crates/socket-patch-core/src/api/client.rs +++ b/crates/socket-patch-core/src/api/client.rs @@ -14,7 +14,7 @@ use crate::api::ranking::severity_order as get_severity_order; use crate::api::ranking::{cmp_batch_infos, cmp_search_results}; use crate::api::retry::{ is_retryable_status, jitter_sample as retry_jitter, parse_retry_after, ApiRetry, - ApiRetryPolicy, RetryHooks, + ApiRetryPolicy, ApiTimeouts, RetryHooks, }; use crate::api::types::*; use crate::api::vendor_prefetch::VendorPrefetch; @@ -49,6 +49,19 @@ fn network_error_detail(e: &reqwest::Error) -> String { msg } +/// `Response::json` reads the body before decoding it. A read timeout is a +/// transport failure, not malformed JSON; retain its cause chain. +fn json_response_error(error: reqwest::Error, context: &str) -> ApiError { + if error.is_timeout() || error.is_body() { + ApiError::Network(format!( + "Network error reading {context} body: {}", + network_error_detail(&error) + )) + } else { + ApiError::Parse(format!("Failed to parse {context}: {error}")) + } +} + /// The readable part of a non-2xx response body, for an error message: the /// `error.message` / `message` / `error` string of a JSON body (the API's /// error shape), otherwise the trimmed body text. Empty when there is @@ -376,28 +389,11 @@ impl ApiClient { /// (User-Agent, Accept, and optionally Authorization). pub fn new(options: ApiClientOptions) -> Self { let api_url = options.api_url.trim_end_matches('/').to_string(); - - let mut default_headers = HeaderMap::new(); - default_headers.insert( - header::USER_AGENT, - HeaderValue::from_static(USER_AGENT_VALUE), - ); - default_headers.insert(header::ACCEPT, HeaderValue::from_static("application/json")); - - if let Some(ref token) = options.api_token { - if let Ok(hv) = HeaderValue::from_str(&format!("Bearer {}", token)) { - default_headers.insert(header::AUTHORIZATION, hv); - } - } - - let client = reqwest::Client::builder() - .default_headers(default_headers) - .build() - .expect("failed to build reqwest client"); + let timeouts = ApiTimeouts::default(); Self { - client, - plain: plain_client(), + client: api_client(options.api_token.as_deref(), &timeouts), + plain: plain_client(&timeouts), api_url, api_token: options.api_token, use_public_proxy: options.use_public_proxy, @@ -420,6 +416,15 @@ impl ApiClient { self } + /// Override the connect and stalled-read bounds of both HTTP clients + /// (tests use short ones). Rebuilds the clients, so call it before + /// cloning the client. + pub fn with_api_timeouts(mut self, timeouts: ApiTimeouts) -> Self { + self.client = api_client(self.api_token.as_deref(), &timeouts); + self.plain = plain_client(&timeouts); + self + } + /// Wait for a [`Self::proxy_batch_slots`] slot; held until dropped. async fn proxy_batch_slot(&self) -> tokio::sync::OwnedSemaphorePermit { Arc::clone(&self.proxy_batch_slots) @@ -611,7 +616,7 @@ impl ApiClient { let body = resp .json::() .await - .map_err(|e| ApiError::Parse(format!("Failed to parse response: {}", e)))?; + .map_err(|e| json_response_error(e, "response"))?; return Ok(Some(body)); } if status == StatusCode::NOT_FOUND { @@ -855,7 +860,7 @@ impl ApiClient { let parsed = resp .json::() .await - .map_err(|e| ApiError::Parse(format!("Failed to parse response: {}", e)))?; + .map_err(|e| json_response_error(e, "response"))?; return Ok(Some(parsed)); } if let Some(err) = classify_auth_error(status, true) { @@ -1543,10 +1548,7 @@ impl ApiClient { // A body cut off (or timed out) mid-transfer is transport, // not a malformed answer. let hint = (e.is_timeout() || e.is_body()).then_some(None); - ( - ApiError::Parse(format!("Failed to parse package response: {e}")), - hint, - ) + (json_response_error(e, "package response"), hint) })?; return Ok(parsed.results); } @@ -1905,18 +1907,42 @@ enum ServeDownload { Failed(ApiError), } +/// Build the authenticated `reqwest::Client` ([`ApiClient`]'s `client` +/// field): User-Agent, `Accept: application/json` and, given a token, the +/// Socket bearer, bounded by `timeouts`. +fn api_client(api_token: Option<&str>, timeouts: &ApiTimeouts) -> reqwest::Client { + let mut default_headers = HeaderMap::new(); + default_headers.insert( + header::USER_AGENT, + HeaderValue::from_static(USER_AGENT_VALUE), + ); + default_headers.insert(header::ACCEPT, HeaderValue::from_static("application/json")); + + if let Some(token) = api_token { + if let Ok(hv) = HeaderValue::from_str(&format!("Bearer {}", token)) { + default_headers.insert(header::AUTHORIZATION, hv); + } + } + + timeouts + .apply(reqwest::Client::builder().default_headers(default_headers)) + .build() + .expect("failed to build reqwest client") +} + /// Build a plain `reqwest::Client` carrying only the User-Agent — no /// Authorization. Built once per [`ApiClient`] (its `plain` field) for the /// public-proxy POST and the grant-tokenized serve GETs, where sending the -/// Socket bearer would leak it to a third party. -fn plain_client() -> reqwest::Client { +/// Socket bearer would leak it to a third party. Bounded by `timeouts`, +/// like the authenticated client. +fn plain_client(timeouts: &ApiTimeouts) -> reqwest::Client { let mut headers = HeaderMap::new(); headers.insert( header::USER_AGENT, HeaderValue::from_static(USER_AGENT_VALUE), ); - reqwest::Client::builder() - .default_headers(headers) + timeouts + .apply(reqwest::Client::builder().default_headers(headers)) .build() .expect("failed to build plain reqwest client") } @@ -5456,6 +5482,78 @@ mod vendor_retry_tests { assert_eq!(posts(&server).await, 2, "the timed-out attempt is retried"); } + /// A response can stall after its successful headers arrive. Keep the + /// transport classification and retry hint while reading package JSON. + #[tokio::test] + async fn stalled_post_json_body_is_network_and_retried() { + use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _}; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let uri = format!("http://{}", listener.local_addr().unwrap()); + let requests = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let count = Arc::clone(&requests); + tokio::spawn(async move { + loop { + let Ok((mut socket, _)) = listener.accept().await else { + return; + }; + let count = Arc::clone(&count); + tokio::spawn(async move { + let mut request = Vec::new(); + let mut buffer = [0u8; 4096]; + while !request.windows(4).any(|w| w == b"\r\n\r\n") { + match socket.read(&mut buffer).await { + Ok(0) | Err(_) => return, + Ok(n) => request.extend_from_slice(&buffer[..n]), + } + } + count.fetch_add(1, Ordering::Relaxed); + if socket + .write_all(b"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: 100\r\n\r\n{") + .await + .is_err() + { + return; + } + while let Ok(n) = socket.read(&mut buffer).await { + if n == 0 { + return; + } + } + }); + } + }); + for proxy in [false, true] { + let before = requests.load(Ordering::Relaxed); + let api = ApiClient::new(ApiClientOptions { + api_url: uri.clone(), + api_token: (!proxy).then(|| "tok".into()), + use_public_proxy: proxy, + org_slug: Some("org".into()), + }) + .with_vendor_retry(VendorRetryPolicy { + attempts: 2, + ..fast() + }) + .with_api_timeouts(ApiTimeouts { + connect: Duration::from_secs(5), + read: Duration::from_millis(100), + }); + let (error, retryable) = tokio::time::timeout( + Duration::from_secs(10), + api.request_vendor_references(&[UUID_A.to_string()], false, None), + ) + .await + .expect("body read must be bounded") + .expect_err("partial JSON body must fail"); + assert!( + matches!(&error, ApiError::Network(msg) if msg.contains("timed out")), + "{error:?}" + ); + assert!(retryable, "body timeout keeps the retry hint"); + assert_eq!(requests.load(Ordering::Relaxed) - before, 2); + } + } + /// The archive GET's response headers are bounded the same way. #[tokio::test] async fn stalled_get_times_out_per_attempt_and_is_retried() { diff --git a/crates/socket-patch-core/src/api/retry.rs b/crates/socket-patch-core/src/api/retry.rs index e9b7ec4c2..643214539 100644 --- a/crates/socket-patch-core/src/api/retry.rs +++ b/crates/socket-patch-core/src/api/retry.rs @@ -127,6 +127,50 @@ impl ApiRetryPolicy { } } +/// Default [`ApiTimeouts::connect`]: TCP connect plus the TLS handshake. +pub const API_CONNECT_TIMEOUT: Duration = Duration::from_secs(10); + +/// Default [`ApiTimeouts::read`]: the longest silence on an open +/// connection (waiting for the response headers, or between body chunks). +pub const API_READ_TIMEOUT: Duration = Duration::from_secs(60); + +/// Transport bounds for every patch-API request: the JSON calls, the +/// blob/diff downloads and the vendoring service, on both of +/// [`crate::api::client::ApiClient`]'s HTTP clients. +/// +/// Neither is a total deadline: `read` restarts after every chunk that +/// arrives, so a large download that keeps streaming is never cut off, +/// while a stalled proxy, load balancer or half-open connection fails the +/// request as `ApiError::Network` instead of hanging the run. A stall is a +/// transport error, so the JSON retry loop does not repeat it; the +/// vendoring service's own per-attempt deadlines and retries still apply +/// on top. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ApiTimeouts { + /// Bound on establishing a connection. + pub connect: Duration, + /// Bound on one read waiting for data. + pub read: Duration, +} + +impl Default for ApiTimeouts { + fn default() -> Self { + Self { + connect: API_CONNECT_TIMEOUT, + read: API_READ_TIMEOUT, + } + } +} + +impl ApiTimeouts { + /// `builder` with these bounds applied. + pub fn apply(&self, builder: reqwest::ClientBuilder) -> reqwest::ClientBuilder { + builder + .connect_timeout(self.connect) + .read_timeout(self.read) + } +} + /// [`API_MAX_RETRIES_ENV`]'s value as a retry count, or `None` to keep the /// default (unset, empty, or not a non-negative integer). fn max_retries_override(raw: Option<&str>) -> Option { diff --git a/crates/socket-patch-core/tests/api_timeout_e2e.rs b/crates/socket-patch-core/tests/api_timeout_e2e.rs new file mode 100644 index 000000000..793f56dfb --- /dev/null +++ b/crates/socket-patch-core/tests/api_timeout_e2e.rs @@ -0,0 +1,283 @@ +//! Transport bounds on the patch API (`socket_patch_core::api::retry:: +//! ApiTimeouts`), end to end against a local TCP server that accepts a +//! connection, reads the request and never answers (#570). +//! +//! Each former unbounded path — the JSON calls (`fetch_patch`, the batch +//! search) and the blob/diff downloads, on the authenticated client and on +//! the plain public-proxy client — must fail as `ApiError::Network` within +//! the client's (shortened) read bound instead of hanging. A body that keeps +//! streaming for longer than the bound must still arrive whole: the bound is +//! on silence, not on the total transfer. + +use std::time::{Duration, Instant}; + +use socket_patch_core::api::client::{ApiClient, ApiClientOptions, ApiError}; +use socket_patch_core::api::retry::{ + ApiRetryPolicy, ApiTimeouts, RetryHooks, API_CONNECT_TIMEOUT, API_READ_TIMEOUT, +}; +use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _}; +use tokio::net::TcpListener; + +const HASH: &str = "abcdef0123456789abcdef0123456789abcdef0123456789abcdef0123456789"; +const UUID: &str = "11111111-2222-4333-8444-555555555555"; + +/// The test bound on one silent read. +const READ: Duration = Duration::from_millis(300); +/// How long a bounded call may take in total before the test calls it a +/// hang (generous: CI runners are slow, an unbounded call never returns). +const GUARD: Duration = Duration::from_secs(20); + +/// A server that accepts every connection, reads what the client sends and +/// never writes a byte. Returns its base URL. +async fn stalled_server() -> String { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let Ok((mut sock, _)) = listener.accept().await else { + return; + }; + tokio::spawn(async move { + let mut buf = [0u8; 4096]; + // Drain the request, then hold the connection open, silent. + while let Ok(n) = sock.read(&mut buf).await { + if n == 0 { + return; + } + } + }); + } + }); + format!("http://{addr}") +} + +/// Send successful JSON headers and a partial body, then keep the socket +/// open. This stalls after `send()` has returned, while `json()` reads. +async fn stalled_json_body_server() -> String { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let Ok((mut sock, _)) = listener.accept().await else { + return; + }; + tokio::spawn(async move { + let mut request = Vec::new(); + let mut buf = [0u8; 4096]; + while !request.windows(4).any(|w| w == b"\r\n\r\n") { + match sock.read(&mut buf).await { + Ok(0) | Err(_) => return, + Ok(n) => request.extend_from_slice(&buf[..n]), + } + } + if sock + .write_all(b"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: 100\r\n\r\n{") + .await + .is_err() + { + return; + } + while let Ok(n) = sock.read(&mut buf).await { + if n == 0 { + return; + } + } + }); + } + }); + format!("http://{addr}") +} + +/// A server that answers `200` with a `total`-byte body, sent in `chunks` +/// pieces `gap` apart (each gap shorter than [`READ`], the sum longer). +async fn trickling_server(total: usize, chunks: usize, gap: Duration) -> String { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { + loop { + let Ok((mut sock, _)) = listener.accept().await else { + return; + }; + tokio::spawn(async move { + let mut req = Vec::new(); + let mut buf = [0u8; 4096]; + while !req.windows(4).any(|w| w == b"\r\n\r\n") { + match sock.read(&mut buf).await { + Ok(0) | Err(_) => return, + Ok(n) => req.extend_from_slice(&buf[..n]), + } + } + let head = format!( + "HTTP/1.1 200 OK\r\ncontent-type: application/octet-stream\r\n\ + content-length: {total}\r\nconnection: close\r\n\r\n" + ); + if sock.write_all(head.as_bytes()).await.is_err() { + return; + } + let piece = total / chunks; + for i in 0..chunks { + tokio::time::sleep(gap).await; + let len = if i + 1 == chunks { + total - piece * i + } else { + piece + }; + if sock.write_all(&vec![b'x'; len]).await.is_err() { + return; + } + } + let _ = sock.flush().await; + }); + } + }); + format!("http://{addr}") +} + +/// A client on `uri` (authenticated, or the public proxy) with the short +/// test read bound and JSON retries off. +fn client(uri: &str, proxy: bool) -> ApiClient { + ApiClient::new(ApiClientOptions { + api_url: uri.to_string(), + api_token: (!proxy).then(|| "tok".to_string()), + use_public_proxy: proxy, + org_slug: (!proxy).then(|| "org".to_string()), + }) + .with_api_retry(ApiRetryPolicy::none(), RetryHooks::default()) + .with_api_timeouts(ApiTimeouts { + connect: Duration::from_secs(5), + read: READ, + }) +} + +/// Run `call`, asserting it fails as `ApiError::Network` after the read +/// bound and well before [`GUARD`]. +async fn assert_stall_is_network( + what: &str, + call: impl std::future::Future>, +) { + let started = Instant::now(); + let result = tokio::time::timeout(GUARD, call) + .await + .unwrap_or_else(|_| panic!("{what} still pending after {GUARD:?}: unbounded")); + let elapsed = started.elapsed(); + match result { + Err(ApiError::Network(msg)) => { + assert!( + elapsed >= READ, + "{what} failed after {elapsed:?}, before the {READ:?} bound: {msg}" + ); + } + other => panic!("{what}: expected ApiError::Network, got {other:?}"), + } +} + +#[tokio::test] +async fn authenticated_calls_fail_as_network_on_a_stalled_server() { + let api = client(&stalled_server().await, false); + assert_stall_is_network("fetch_patch", api.fetch_patch(UUID)).await; + assert_stall_is_network( + "search_patches_batch", + api.search_patches_batch(&["pkg:npm/left-pad@1.3.0".to_string()]), + ) + .await; + assert_stall_is_network("fetch_blob", api.fetch_blob(HASH)).await; + assert_stall_is_network("fetch_diff", api.fetch_diff(UUID)).await; +} + +#[tokio::test] +async fn public_proxy_calls_fail_as_network_on_a_stalled_server() { + // The proxy client sends on the plain (header-free) client, so this + // covers the second reqwest client. + let api = client(&stalled_server().await, true); + assert_stall_is_network("fetch_patch", api.fetch_patch(UUID)).await; + assert_stall_is_network( + "search_patches_batch", + api.search_patches_batch(&["pkg:npm/left-pad@1.3.0".to_string()]), + ) + .await; + assert_stall_is_network("fetch_blob", api.fetch_blob(HASH)).await; + assert_stall_is_network("fetch_diff", api.fetch_diff(UUID)).await; +} + +#[tokio::test] +async fn stalled_json_bodies_are_network_errors_on_both_clients() { + let uri = stalled_json_body_server().await; + let mut errors = Vec::new(); + for proxy in [false, true] { + let api = client(&uri, proxy); + let patch = tokio::time::timeout(GUARD, api.fetch_patch(UUID)) + .await + .expect("stalled patch body must time out") + .expect_err("partial patch body must fail"); + errors.push((format!("fetch_patch proxy={proxy}"), patch)); + let batch = tokio::time::timeout( + GUARD, + api.search_patches_batch(&["pkg:npm/left-pad@1.3.0".to_string()]), + ) + .await + .expect("stalled batch body must time out") + .expect_err("partial batch body must fail"); + errors.push((format!("search_patches_batch proxy={proxy}"), batch)); + } + assert!( + errors.iter().all(|(_, error)| matches!( + error, + ApiError::Network(message) if message.contains("timed out") + )), + "stalled JSON bodies must retain the timeout cause as network errors: {errors:?}" + ); +} + +#[tokio::test] +async fn completed_malformed_json_remains_a_parse_error() { + let server = wiremock::MockServer::start().await; + wiremock::Mock::given(wiremock::matchers::any()) + .respond_with(wiremock::ResponseTemplate::new(200).set_body_string("not json")) + .mount(&server) + .await; + for proxy in [false, true] { + let api = client(&server.uri(), proxy); + assert!(matches!( + api.fetch_patch(UUID).await, + Err(ApiError::Parse(_)) + )); + assert!(matches!( + api.search_patches_batch(&["pkg:npm/left-pad@1.3.0".to_string()]) + .await, + Err(ApiError::Parse(_)) + )); + } +} + +#[tokio::test] +async fn a_body_that_keeps_streaming_past_the_bound_still_arrives() { + // 8 chunks 150 ms apart: every silence is under the 300 ms bound, the + // whole transfer (~1.2 s) is four times it. + let total = 64 * 1024; + for proxy in [false, true] { + let api = client( + &trickling_server(total, 8, Duration::from_millis(150)).await, + proxy, + ); + let started = Instant::now(); + let body = tokio::time::timeout(GUARD, api.fetch_blob(HASH)) + .await + .expect("trickled blob still pending") + .expect("a streaming body must not time out") + .expect("200 is a blob"); + assert_eq!(body.len(), total, "proxy={proxy}"); + assert!( + started.elapsed() > READ * 3, + "proxy={proxy}: the body trickled" + ); + } +} + +#[test] +fn default_bounds_are_the_named_constants() { + let t = ApiTimeouts::default(); + assert_eq!(t.connect, API_CONNECT_TIMEOUT); + assert_eq!(t.read, API_READ_TIMEOUT); + assert_eq!(API_CONNECT_TIMEOUT, Duration::from_secs(10)); + assert_eq!(API_READ_TIMEOUT, Duration::from_secs(60)); +}