Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
8 changes: 7 additions & 1 deletion crates/socket-patch-bench/src/fixtures/npm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -623,7 +623,13 @@ pub fn build_yarn_berry(t: &mut Tree, size: Size) -> std::io::Result<Fixture> {
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 ────────────────────────────────────────────────────────────────
Expand Down
160 changes: 129 additions & 31 deletions crates/socket-patch-core/src/api/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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)
Expand Down Expand Up @@ -611,7 +616,7 @@ impl ApiClient {
let body = resp
.json::<T>()
.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 {
Expand Down Expand Up @@ -855,7 +860,7 @@ impl ApiClient {
let parsed = resp
.json::<BatchSearchResponse>()
.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) {
Expand Down Expand Up @@ -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);
}
Expand Down Expand Up @@ -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")
}
Expand Down Expand Up @@ -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() {
Expand Down
44 changes: 44 additions & 0 deletions crates/socket-patch-core/src/api/retry.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<u32> {
Expand Down
Loading
Loading