diff --git a/remote_execution/oss/re_grpc/src/client.rs b/remote_execution/oss/re_grpc/src/client.rs index ecb4443eef8f8..657bd1dd55632 100644 --- a/remote_execution/oss/re_grpc/src/client.rs +++ b/remote_execution/oss/re_grpc/src/client.rs @@ -161,7 +161,8 @@ pub struct RECapabilities { /// Largest size of a message before being uploaded using bytestream service. /// 0 indicates no limit beyond constraint of underlying transport (which is unknown). max_total_batch_size: usize, - /// Compressors supported by the "compressed-blobs" bytestream resources. + /// Compressors supported by the "compressed-blobs" bytestream resources. Also used to decide + /// whether we can request zstd-compressed `BatchReadBlobs` responses. supported_compressors: Vec, } @@ -175,6 +176,8 @@ pub struct RERuntimeOpts { cas_ttl_secs: i64, /// Maximum number of digests per `FindMissingBlobs` RPC. find_missing_blobs_batch_size: usize, + /// Whether to request zstd-compressed `BatchReadBlobs` responses (server supports zstd). + batch_read_zstd: bool, } struct InstanceName(Option); @@ -225,6 +228,15 @@ impl Compressor { } } +/// Decompress a zstd-compressed blob returned in a `BatchReadBlobs` response. +async fn zstd_decompress_blob(data: &[u8]) -> anyhow::Result> { + let mut decoder = ZstdDecoder::new(Cursor::new(data)); + decoder.multiple_members(true); + let mut out = Vec::new(); + decoder.read_to_end(&mut out).await?; + Ok(out) +} + pub struct REClientBuilder; impl REClientBuilder { @@ -321,6 +333,9 @@ impl REClientBuilder { // on the TTL of the remote blob. cas_ttl_secs: opts.cas_ttl_secs.unwrap_or(3 * 60 * 60), find_missing_blobs_batch_size: opts.find_missing_blobs_batch_size.unwrap_or(100), + batch_read_zstd: capabilities + .supported_compressors + .contains(&Compressor::Zstd), }, capabilities, instance_name, @@ -908,6 +923,7 @@ impl REClient { request: DownloadRequest, ) -> anyhow::Result { download_impl( + &self.runtime_opts, &self.instance_name, request, self.bystream_compressor, @@ -1276,6 +1292,7 @@ fn convert_t_action_result2(t_action_result: TActionResult2) -> anyhow::Result( + opts: &RERuntimeOpts, instance_name: &InstanceName, request: DownloadRequest, bystream_compressor: Option, @@ -1352,6 +1369,17 @@ where let inlined_digests = request.inlined_digests.unwrap_or_default(); let file_digests = request.file_digests.unwrap_or_default(); + // Advertise zstd for BatchReadBlobs responses when the server supports it; Identity is always + // acceptable as a fallback and the server decides per-blob what to actually return. + let acceptable_compressors = if opts.batch_read_zstd { + vec![ + compressor::Value::Zstd as i32, + compressor::Value::Identity as i32, + ] + } else { + vec![compressor::Value::Identity as i32] + }; + let mut curr_size = 0; let mut requests = vec![]; let mut curr_digests = vec![]; @@ -1372,7 +1400,7 @@ where let read_blob_req = BatchReadBlobsRequest { instance_name: instance_name.as_str().to_owned(), digests: std::mem::take(&mut curr_digests), - acceptable_compressors: vec![compressor::Value::Identity as i32], + acceptable_compressors: acceptable_compressors.clone(), ..Default::default() }; requests.push(read_blob_req); @@ -1385,7 +1413,7 @@ where let read_blob_req = BatchReadBlobsRequest { instance_name: instance_name.as_str().to_owned(), digests: std::mem::take(&mut curr_digests), - acceptable_compressors: vec![compressor::Value::Identity as i32], + acceptable_compressors: acceptable_compressors.clone(), ..Default::default() }; requests.push(read_blob_req); @@ -1402,7 +1430,20 @@ where for r in resp.responses.into_iter() { let digest = tdigest_from(r.digest.context("Response digest not found.")?); check_status(r.status.unwrap_or_default())?; - batched_blobs_response.insert(digest, r.data); + let data = if r.compressor == compressor::Value::Identity as i32 { + r.data + } else if r.compressor == compressor::Value::Zstd as i32 { + zstd_decompress_blob(&r.data) + .await + .context("Failed to decompress zstd BatchReadBlobs response")? + } else { + return Err(anyhow::anyhow!( + "BatchReadBlobs response for `{}` used an unsupported compressor: {}", + digest, + r.compressor + )); + }; + batched_blobs_response.insert(digest, data); } } @@ -1844,6 +1885,16 @@ mod tests { use super::*; + pub(super) fn test_re_runtime_opts() -> RERuntimeOpts { + RERuntimeOpts { + use_fbcode_metadata: false, + max_concurrent_uploads_per_action: None, + cas_ttl_secs: 0, + find_missing_blobs_batch_size: 100, + batch_read_zstd: false, + } + } + #[tokio::test] async fn test_download_named() -> anyhow::Result<()> { let work = tempfile::tempdir()?; @@ -1907,6 +1958,7 @@ mod tests { }; download_impl( + &test_re_runtime_opts(), &InstanceName(None), req, None, @@ -2014,6 +2066,7 @@ mod tests { }; download_impl( + &test_re_runtime_opts(), &InstanceName(None), req, None, @@ -2096,6 +2149,7 @@ mod tests { }; let res = download_impl( + &test_re_runtime_opts(), &InstanceName(None), req, None, @@ -2183,6 +2237,7 @@ mod tests { let counter = AtomicU16::new(0); let res = download_impl( + &test_re_runtime_opts(), &InstanceName(None), req, None, @@ -2252,6 +2307,7 @@ mod tests { }; let res = download_impl( + &test_re_runtime_opts(), &InstanceName(None), req, None, @@ -2308,6 +2364,7 @@ mod tests { let res = BatchReadBlobsResponse { responses: vec![] }; let res = download_impl( + &test_re_runtime_opts(), &InstanceName(None), req, None, @@ -2347,6 +2404,7 @@ mod tests { }; download_impl( + &test_re_runtime_opts(), &InstanceName(Some("instance".to_owned())), req, None, @@ -2442,6 +2500,65 @@ mod tests { Ok(()) } + #[tokio::test] + async fn test_batch_download_zstd() -> anyhow::Result<()> { + // The server returns a zstd-compressed blob; we advertise zstd and decompress it. + let blob = vec![9u8; 500]; + let digest = TDigest { + hash: "dd".to_owned(), + size_in_bytes: blob.len() as i64, + ..Default::default() + }; + + let mut compressed = vec![]; + ZstdEncoder::new(Cursor::new(blob.clone())) + .read_to_end(&mut compressed) + .await + .unwrap(); + + let req = DownloadRequest { + inlined_digests: Some(vec![digest.clone()]), + ..Default::default() + }; + + let res = BatchReadBlobsResponse { + responses: vec![batch_read_blobs_response::Response { + digest: Some(tdigest_to(digest.clone())), + data: compressed.clone(), + status: Some(Status::default()), + compressor: compressor::Value::Zstd as i32, + }], + }; + + let mut opts = test_re_runtime_opts(); + opts.batch_read_zstd = true; + + let expected = blob.clone(); + let d_resp = download_impl( + &opts, + &InstanceName(None), + req, + None, + 10000, + |req| { + let res = res.clone(); + async move { + // We advertised zstd as an acceptable response compressor. + assert!( + req.acceptable_compressors + .contains(&(compressor::Value::Zstd as i32)) + ); + Ok(res) + } + }, + |_digest| async move { anyhow::Ok(Box::pin(futures::stream::iter(vec![]))) }, + ) + .await?; + + assert_eq!(d_resp.inlined_blobs.unwrap()[0].blob, expected); + Ok(()) + } + #[tokio::test] async fn test_upload_large_named() -> anyhow::Result<()> { let blob_data = vec![ @@ -2975,6 +3092,7 @@ async fn test_download_compressed() -> anyhow::Result<()> { let compressed_data_ref = &compressed_data; let d_resp = download_impl( + &tests::test_re_runtime_opts(), &InstanceName(None), DownloadRequest { inlined_digests: Some(vec![TDigest {