Skip to content
Open
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
28 changes: 9 additions & 19 deletions crates/socket-patch-core/src/patch/jvm_jar.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,8 +25,6 @@
use std::collections::HashMap;
use std::path::{Path, PathBuf};

use sha1::Digest as _;

use crate::crawlers::gradle_cache;
use crate::hash::git_sha256::compute_git_sha256_from_bytes;
use crate::manifest::schema::PatchFileInfo;
Expand All @@ -36,6 +34,7 @@ use crate::patch::apply::{
};
use crate::patch::rollback::{RollbackResult, VerifyRollbackResult, VerifyRollbackStatus};
use crate::patch::sidecars::{self, maven as maven_sidecars};
use crate::utils::digest::{sha1_hex_of, sha256_hex_of};
use crate::utils::purl::{parse_maven_purl, purl_qualifier};
use crate::vendor::VendorServiceConfig;

Expand Down Expand Up @@ -270,7 +269,7 @@ pub fn derived_copies_in(
// reader fails fast on a FIFO or device instead of wedging in open(2).
let current = crate::utils::fs::read_regular_to_bytes_sync(&jar)
.ok()
.map(|b| sha1_hex(&b));
.map(|b| sha1_hex_of(&b));
let written = std::fs::metadata(&jar).and_then(|m| m.modified()).ok();
let unverified = found
.unknown
Expand All @@ -279,7 +278,7 @@ pub fn derived_copies_in(
let Ok(copy) = crate::utils::fs::read_regular_to_bytes_sync(p) else {
return true;
};
if Some(sha1_hex(&copy)) == current {
if Some(sha1_hex_of(&copy)) == current {
return false;
}
let made = std::fs::metadata(p).and_then(|m| m.modified()).ok();
Expand Down Expand Up @@ -352,20 +351,11 @@ fn unpatched_members(
Ok(members)
}

fn sha256_hex(bytes: &[u8]) -> String {
use sha2::Digest as _;
hex::encode(sha2::Sha256::digest(bytes))
}

fn sha1_hex(bytes: &[u8]) -> String {
hex::encode(sha1::Sha1::digest(bytes))
}

/// `<socket_dir>/jvm-originals/<sha256>.jar`.
pub fn backup_path(socket_dir: &Path, original: &[u8]) -> PathBuf {
socket_dir
.join(ORIGINALS_DIR)
.join(format!("{}.jar", sha256_hex(original)))
.join(format!("{}.jar", sha256_hex_of(original)))
}

/// Keep `original` under [`ORIGINALS_DIR`] (content-addressed: an existing
Expand Down Expand Up @@ -625,7 +615,7 @@ async fn find_backup(restore: &JarRestore<'_>, dir: &Path, current: &[u8]) -> Op
continue;
};
if let Some(hash) = &gradle_hash {
if !gradle_cache::hash_eq(hash, &sha1_hex(&bytes)) {
if !gradle_cache::hash_eq(hash, &sha1_hex_of(&bytes)) {
continue;
}
}
Expand Down Expand Up @@ -663,7 +653,7 @@ async fn upstream_for_gradle_copy(purl: &str, jar_leaf: &str, dir: &Path) -> Opt
)
.await
.ok()?;
gradle_cache::hash_eq(hash, &sha1_hex(&bytes)).then_some(bytes)
gradle_cache::hash_eq(hash, &sha1_hex_of(&bytes)).then_some(bytes)
}

fn rollback_result(purl: &str, dir: &Path) -> RollbackResult {
Expand Down Expand Up @@ -1066,7 +1056,7 @@ mod tests {
std::fs::create_dir_all(&copy).unwrap();
let original = pristine_jar();
std::fs::write(copy.join("lib-1.0.jar"), &original).unwrap();
let sha1_text = format!("{}\n", sha1_hex(&original));
let sha1_text = format!("{}\n", sha1_hex_of(&original));
std::fs::write(copy.join("lib-1.0.jar.sha1"), &sha1_text).unwrap();

let files = record();
Expand All @@ -1082,7 +1072,7 @@ mod tests {
);
assert_eq!(
std::fs::read_to_string(copy.join("lib-1.0.jar.sha1")).unwrap(),
format!("{}\n", sha1_hex(&service))
format!("{}\n", sha1_hex_of(&service))
);

let restore = JarRestore {
Expand Down Expand Up @@ -1123,7 +1113,7 @@ mod tests {
let version = d
.path()
.join(".gradle/caches/modules-2/files-2.1/com.example/lib/1.0");
let hash_dir = version.join(sha1_hex(&original));
let hash_dir = version.join(sha1_hex_of(&original));
std::fs::create_dir_all(&hash_dir).unwrap();
std::fs::write(hash_dir.join("lib-1.0.jar"), patched_jar()).unwrap();

Expand Down
4 changes: 1 addition & 3 deletions crates/socket-patch-core/src/patch/sidecars/maven.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,6 @@

use std::path::{Path, PathBuf};

use sha1::Digest as _;

use super::{
SidecarAdvisory, SidecarAdvisoryCode, SidecarError, SidecarFile, SidecarFileAction,
SidecarPayload, SidecarSeverity,
Expand All @@ -44,7 +42,7 @@ impl Algo {

fn digest(self, bytes: &[u8]) -> String {
match self {
Algo::Sha1 => hex::encode(sha1::Sha1::digest(bytes)),
Algo::Sha1 => crate::utils::digest::sha1_hex_of(bytes),
Algo::Md5 => hex::encode(md5(bytes)),
}
}
Expand Down
1 change: 1 addition & 0 deletions crates/socket-patch-core/src/utils/digest.rs
Original file line number Diff line number Diff line change
Expand Up @@ -135,6 +135,7 @@ mod tests {
/// when you move it onto the helpers above; the test fails on a stale
/// entry as well as on a new inline copy.
const PENDING_INLINE_DIGESTS: &[&str] = &[
"crawlers/gradle_cache.rs",
"utils/group_commit.rs",
"vendor/jvm/mod.rs",
"vendor/maven_repo.rs",
Expand Down
5 changes: 1 addition & 4 deletions crates/socket-patch-core/src/vendor/maven_repo.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::time::Duration;

use serde_json::Value;
use sha1::Sha1;
Expand Down Expand Up @@ -1759,9 +1758,7 @@ async fn fetch_pom_bytes(url: &str) -> Result<Vec<u8>, String> {
}

pub(crate) async fn fetch_registry_bytes(url: &str, cap: u64) -> Result<Vec<u8>, String> {
let client = reqwest::Client::builder()
.user_agent(MAVEN_USER_AGENT)
.timeout(Duration::from_secs(60))
let client = super::registry_fetch::registry_client_builder(MAVEN_USER_AGENT)
.build()
.map_err(|e| format!("build http client: {e}"))?;
let resp = client
Expand Down
182 changes: 153 additions & 29 deletions crates/socket-patch-core/src/vendor/registry_fetch.rs
Original file line number Diff line number Diff line change
@@ -1,12 +1,12 @@
//! Bounded archive readers, integrity verification and registry metadata transport.

use std::path::{Path, PathBuf};
use std::time::Duration;

use base64::Engine as _;
use sha1::Sha1;
use sha2::{Digest, Sha256, Sha384, Sha512};

use crate::api::retry::ApiTimeouts;
use crate::constants::USER_AGENT;
use crate::patch::apply::is_safe_relative_subpath;

Expand Down Expand Up @@ -43,13 +43,61 @@ pub enum FetchError {
pub type RegistryClient = reqwest::Client;

pub fn build_registry_client() -> RegistryClient {
reqwest::Client::builder()
.user_agent(USER_AGENT)
.timeout(Duration::from_secs(60))
registry_client_builder(USER_AGENT)
.build()
.unwrap_or_else(|_| reqwest::Client::new())
}

/// The one builder behind every registry client (npm-family, PyPI, Go,
/// NuGet, Maven), sending `user_agent`. It applies the shared
/// [`ApiTimeouts`] transport policy: a connect bound plus an idle read
/// bound that restarts on every chunk, and no total deadline, so a large
/// artifact that keeps streaming is never cut off while a stalled host
/// still fails the fetch.
pub(crate) fn registry_client_builder(user_agent: &str) -> reqwest::ClientBuilder {
registry_timeouts().apply(reqwest::Client::builder().user_agent(user_agent))
}

fn registry_timeouts() -> ApiTimeouts {
#[cfg(test)]
if let Some(t) = test_timeouts::get() {
return t;
}
ApiTimeouts::default()
}

/// Test-only override of [`registry_timeouts`] for the current thread, so a
/// test can prove the idle bound and the absence of a total deadline in
/// seconds rather than minutes. `#[tokio::test]` runs on one thread.
#[cfg(test)]
pub(crate) mod test_timeouts {
use std::cell::Cell;

use crate::api::retry::ApiTimeouts;

thread_local! {
static OVERRIDE: Cell<Option<ApiTimeouts>> = const { Cell::new(None) };
}

pub(crate) fn get() -> Option<ApiTimeouts> {
OVERRIDE.with(Cell::get)
}

/// Shortens the bounds until the returned guard drops.
pub(crate) fn set(t: ApiTimeouts) -> Guard {
OVERRIDE.with(|c| c.set(Some(t)));
Guard
}

pub(crate) struct Guard;

impl Drop for Guard {
fn drop(&mut self) {
OVERRIDE.with(|c| c.set(None));
}
}
}

/// The npm registry base after the env override.
pub fn npm_registry_base() -> String {
std::env::var("SOCKET_NPM_REGISTRY")
Expand Down Expand Up @@ -1170,14 +1218,14 @@ fn walk_zip_with_prefix(
Ok(())
}

/// Capped download. http(s) only; the cap is enforced on the declared
/// Content-Length AND the actual stream (a lying server cannot blow past
/// it).
/// Capped download. http(s) only; [`crate::utils::http::read_capped`]
/// enforces [`MAX_DOWNLOAD_BYTES`] on the declared Content-Length AND the
/// actual stream (a lying server cannot blow past it).
pub(crate) async fn download(client: &reqwest::Client, url: &str) -> Result<Vec<u8>, String> {
if !(url.starts_with("https://") || url.starts_with("http://")) {
return Err(format!("refusing non-http(s) artifact URL `{url}`"));
}
let mut resp = client
let resp = client
.get(url)
.send()
.await
Expand All @@ -1186,27 +1234,9 @@ pub(crate) async fn download(client: &reqwest::Client, url: &str) -> Result<Vec<
if !status.is_success() {
return Err(format!("GET {url}: HTTP {status}"));
}
if let Some(len) = resp.content_length() {
if len > MAX_DOWNLOAD_BYTES {
return Err(format!(
"{url}: artifact is {len} bytes (cap {MAX_DOWNLOAD_BYTES})"
));
}
}
let mut bytes: Vec<u8> = Vec::new();
while let Some(chunk) = resp
.chunk()
crate::utils::http::read_capped(resp, MAX_DOWNLOAD_BYTES, "registry artifact")
.await
.map_err(|e| format!("reading {url}: {e}"))?
{
if bytes.len() as u64 + chunk.len() as u64 > MAX_DOWNLOAD_BYTES {
return Err(format!(
"{url}: artifact exceeds the {MAX_DOWNLOAD_BYTES}-byte cap"
));
}
bytes.extend_from_slice(&chunk);
}
Ok(bytes)
.map_err(|e| format!("{url}: {e}"))
}

/// Verify archive bytes against lock-recorded integrity. Berry cache checksums
Expand Down Expand Up @@ -2993,12 +3023,106 @@ mod tests {
.await
.unwrap_err();
assert!(
err.contains("exceeds the") && err.contains("cap"),
err.contains("exceeded") && err.contains("cap"),
"the stream cap must fire without a Content-Length: {err}"
);
server.abort();
}

/// Serves one GET per accepted connection: a 200 head declaring
/// `chunks × 1 KiB`, then each 1 KiB chunk after its `gaps` delay.
async fn paced_server(
gaps: Vec<std::time::Duration>,
) -> (std::net::SocketAddr, tokio::task::JoinHandle<()>) {
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
loop {
let Ok((mut sock, _)) = listener.accept().await else {
return;
};
let gaps = gaps.clone();
tokio::spawn(async move {
let mut buf = [0u8; 4096];
let _ = sock.read(&mut buf).await; // request head
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
gaps.len() * 1024
);
if sock.write_all(head.as_bytes()).await.is_err() {
return;
}
for gap in gaps {
tokio::time::sleep(gap).await;
if sock.write_all(&[b'x'; 1024]).await.is_err() {
return;
}
}
});
}
});
(addr, server)
}

/// Every registry client: hosted upstream restore's
/// `build_registry_client` + `download`, and vendored Maven's
/// `fetch_registry_bytes` (which sends Maven's own user agent).
async fn fetch_through_every_registry_client(url: &str) -> Vec<Result<Vec<u8>, String>> {
vec![
download(&build_registry_client(), url).await,
crate::vendor::maven_repo::fetch_registry_bytes(url, MAX_DOWNLOAD_BYTES).await,
]
}

const SHORT_BOUNDS: ApiTimeouts = ApiTimeouts {
connect: std::time::Duration::from_secs(2),
read: std::time::Duration::from_millis(400),
};

#[tokio::test]
async fn registry_clients_have_no_total_deadline() {
// #872: the registry clients set a 60 s whole-request deadline, so
// a slow but steady download was aborted mid-body. Under the shared
// `ApiTimeouts` policy only silence counts: a body that trickles
// for 4× the (shortened) idle bound, never pausing that long,
// arrives whole through every registry client.
let _bounds = test_timeouts::set(SHORT_BOUNDS);
let gaps = vec![std::time::Duration::from_millis(100); 16];
let (addr, server) = paced_server(gaps).await;
let url = format!("http://{addr}/slow.tgz");
for got in fetch_through_every_registry_client(&url).await {
assert_eq!(
got.expect("a progressing body must not time out").len(),
16 * 1024
);
}
server.abort();
}

#[tokio::test]
async fn registry_clients_fail_a_body_that_stalls_past_the_idle_bound() {
// A connection that goes silent mid-body fails at the idle bound
// instead of holding the run until the server resumes (here 5 s
// later; on `main` both clients waited it out and succeeded).
let _bounds = test_timeouts::set(SHORT_BOUNDS);
let mut gaps = vec![std::time::Duration::ZERO; 4];
gaps.push(std::time::Duration::from_secs(5));
let (addr, server) = paced_server(gaps).await;
let url = format!("http://{addr}/stall.tgz");
let started = std::time::Instant::now();
for got in fetch_through_every_registry_client(&url).await {
let err = got.expect_err("a stalled body must fail at the idle bound");
assert!(err.contains("error reading"), "{err}");
}
assert!(
started.elapsed() < std::time::Duration::from_secs(4),
"both fetches must give up at the idle bound, took {:?}",
started.elapsed()
);
server.abort();
}

#[test]
fn total_decompressed_cap_fails_closed_across_zip_extractors() {
// The per-entry actual-bytes guards mean only HONEST content reaches
Expand Down
Loading