use crate::uploads::{ConnectionUploads, UploadQueue, UploadSink};
use crate::{
ActionPrediction, ActionPromiseCompletion, ActionPromiseState, BlobPackLimits, CacheDigest,
CacheDirectory, LocalActionCache, LocalCas, MAX_STAGED_BLOB_PACK_BYTES,
MAX_STAGED_BLOB_PACK_ITEMS, ManifestPutOutcome, RemoteActionResult, RemoteCacheClient,
RemoteCacheMode, RustcMetadata, TaskActionManifest, blob_pack_chunk, canonical_json,
};
use eyre::{Context, Result, bail};
use futures_util::{FutureExt, StreamExt, future::BoxFuture, stream};
use log::warn;
use std::collections::{BTreeMap, BTreeSet};
use std::fs;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, Weak};
use std::time::{Duration, Instant, SystemTime};
const ACTION_PROMISE_WAIT: Duration = Duration::from_secs(60 * 60);
use tokio::io::{AsyncBufRead, AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
mod file_digest;
mod manifest;
mod prefetch;
mod stats;
mod wire;
#[cfg(test)]
pub(crate) use prefetch::select_prefetch_actions;
pub use file_digest::{
FileDigestCache, FileDigestScope, FileIdentity, NoFileDigestCache, RecordedFileDigest,
};
pub use manifest::{is_task_identity, task_manifest_actions};
use manifest::{
merge_remote_task_manifest, merge_task_manifests, task_manifest_dir, validate_task_identity,
validate_task_manifest,
};
pub use stats::{AgentStats, CompilerStats};
use wire::MAX_REQUEST_BYTES;
pub use wire::{
AGENT_PROTOCOL_VERSION, AgentEvent, AgentEventObserver, AgentRequest, AgentResponse,
RestoreStats,
};
const MAX_EXECUTABLE_IDENTITIES: usize = 64;
const MAX_EXECUTABLE_IDENTITY_SIZE: usize = 64 * 1024;
const MAX_EXECUTABLE_IDENTITY_BYTES: usize = 256 * 1024;
const TASK_ACTION_MANIFEST_VERSION: u8 = 1;
const MAX_TASK_ACTION_PREDICTIONS: usize = 16 * 1024;
const MAX_WARNING_BYTES: usize = 4 * 1024;
const MAX_WARNINGS: usize = 128;
const MAX_FILE_DIGEST_BATCH: usize = 16 * 1024;
const MAX_FILE_DIGEST_ENTRIES: usize = 1024 * 1024;
const MAX_REMOTE_TRANSFERS: usize = 64;
const MAX_PREFETCH_TRANSFERS: usize = 48;
const MAX_PREFETCH_ACTIONS: usize = 256;
const MAX_PREFETCH_ACTION_WAVE: usize = 32;
const MAX_PREFETCH_BATCH_LOOKUPS: usize = 1;
const PREFETCH_ACTION_BATCH_DELAY: Duration = Duration::from_millis(5);
const MAX_PREFETCH_DIRECTORY_OBJECTS: usize = 100_000;
const MAX_PREFETCH_OBJECTS_PER_WAVE: usize = 100_000;
const DEFAULT_MAX_REMOTE_DOWNLOAD_BYTES: u64 = 5 * 1024 * 1024 * 1024;
pub struct AgentRemoteCache {
pub client: RemoteCacheClient,
pub mode: RemoteCacheMode,
pub staging_dir: PathBuf,
}
#[derive(Default)]
struct AtomicAgentStats {
lookups: AtomicU64,
unconsulted: AtomicU64,
hits: AtomicU64,
stores: AtomicU64,
stored_bytes: AtomicU64,
verifications: AtomicU64,
divergences: AtomicU64,
downloaded_bytes: AtomicU64,
uploaded_bytes: AtomicU64,
background_uploads: AtomicU64,
background_upload_failures: AtomicU64,
remote_blob_pack_uploads: AtomicU64,
remote_blob_pack_upload_blobs: AtomicU64,
upload_drain_duration_ns: AtomicU64,
prefetched_actions: AtomicU64,
predictions_loaded: AtomicU64,
remote_failures: AtomicU64,
remote_manifest_lookups: AtomicU64,
remote_manifest_lookup_duration_ns: AtomicU64,
remote_action_lookups: AtomicU64,
remote_action_lookup_duration_ns: AtomicU64,
remote_blob_requests: AtomicU64,
remote_blob_pack_requests: AtomicU64,
remote_blob_pack_blobs: AtomicU64,
remote_blob_transfer_duration_ns: AtomicU64,
local_cas_write_duration_ns: AtomicU64,
prefetch_runs: AtomicU64,
prefetch_duration_ns: AtomicU64,
materialization_duration_ns: AtomicU64,
bypasses: Mutex<BTreeMap<String, u64>>,
avoided_compiler_duration_ns: AtomicU64,
compiler: Mutex<BTreeMap<String, CompilerStats>>,
slow_compilations: Mutex<BTreeMap<String, u64>>,
restored_output_files: AtomicU64,
restored_output_bytes: AtomicU64,
reflinked_output_files: AtomicU64,
reflinked_output_bytes: AtomicU64,
copied_output_files: AtomicU64,
copied_output_bytes: AtomicU64,
reused_output_files: AtomicU64,
reused_output_bytes: AtomicU64,
}
struct AtomicDurationTimer<'a> {
started: Instant,
target: &'a AtomicU64,
}
impl<'a> AtomicDurationTimer<'a> {
fn start(target: &'a AtomicU64) -> Self {
Self {
started: Instant::now(),
target,
}
}
}
impl Drop for AtomicDurationTimer<'_> {
fn drop(&mut self) {
atomic_saturating_add(self.target, duration_ns(self.started));
}
}
fn duration_ns(started: Instant) -> u64 {
started.elapsed().as_nanos().try_into().unwrap_or(u64::MAX)
}
fn validate_crate_name(crate_name: Option<&str>) -> Result<()> {
if let Some(crate_name) = crate_name
&& (crate_name.len() > 256 || crate_name.contains(['\0', '\n', '\r']))
{
bail!("invalid compiler crate name");
}
Ok(())
}
fn atomic_saturating_add(target: &AtomicU64, value: u64) {
let _ = target.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
Some(current.saturating_add(value))
});
}
fn queue_prefetch_digest(
verified: &BTreeMap<CacheDigest, PathBuf>,
pending: &mut BTreeMap<CacheDigest, ()>,
digest: CacheDigest,
) {
if verified.contains_key(&digest) || pending.contains_key(&digest) {
return;
}
pending.insert(digest, ());
}
fn queue_prefetch_directory(
seen: &BTreeMap<CacheDigest, ()>,
pending: &mut BTreeMap<CacheDigest, ()>,
digest: CacheDigest,
limit: usize,
) -> bool {
if seen.contains_key(&digest) || pending.contains_key(&digest) {
return true;
}
if seen.len().saturating_add(pending.len()) >= limit {
return false;
}
pending.insert(digest, ());
true
}
#[derive(Clone)]
pub struct CacheAgent {
cas: LocalCas,
actions: LocalActionCache,
verified_blobs: Arc<Mutex<BTreeMap<CacheDigest, VerifiedBlob>>>,
version: Arc<str>,
write_locks: Arc<Mutex<BTreeMap<CacheDigest, Weak<tokio::sync::Mutex<()>>>>>,
action_locks: Arc<Mutex<BTreeMap<CacheDigest, Weak<tokio::sync::Mutex<()>>>>>,
stats: Arc<AtomicAgentStats>,
observer: Option<Arc<dyn AgentEventObserver>>,
executable_identities: Arc<Mutex<BTreeMap<ExecutableIdentityKey, Vec<u8>>>>,
manifest_dir: Arc<PathBuf>,
task_actions: Arc<Mutex<BTreeMap<String, TaskActionState>>>,
next_task_run: Arc<AtomicU64>,
manifest_write_lock: Arc<Mutex<()>>,
remote: Option<Arc<RemoteCacheClient>>,
remote_mode: RemoteCacheMode,
remote_staging_dir: Arc<PathBuf>,
remote_download_limit: u64,
remote_download_bytes: Arc<AtomicU64>,
pending_remote_actions: Arc<Mutex<BTreeMap<CacheDigest, RemoteActionResult>>>,
remote_transfers: Arc<tokio::sync::Semaphore>,
prefetch_transfers: Arc<tokio::sync::Semaphore>,
prefetch_tasks: Arc<Mutex<Vec<tokio::task::JoinHandle<()>>>>,
warnings: Arc<Mutex<BTreeSet<String>>>,
file_digests: Arc<Mutex<BTreeMap<(FileDigestScope, PathBuf), RecordedFileDigest>>>,
uploads: Option<UploadQueue>,
}
struct AgentUploadSink {
stats: Arc<AtomicAgentStats>,
}
impl UploadSink for AgentUploadSink {
fn record_blob_uploaded(&self, bytes: u64) {
self.stats
.background_uploads
.fetch_add(1, Ordering::Relaxed);
self.stats
.uploaded_bytes
.fetch_add(bytes, Ordering::Relaxed);
}
fn record_action_uploaded(&self) {
self.stats
.background_uploads
.fetch_add(1, Ordering::Relaxed);
}
fn record_blob_pack_uploaded(&self, blobs: u64) {
self.stats
.remote_blob_pack_uploads
.fetch_add(1, Ordering::Relaxed);
self.stats
.remote_blob_pack_upload_blobs
.fetch_add(blobs, Ordering::Relaxed);
}
fn record_upload_failure(&self) {
self.stats
.background_upload_failures
.fetch_add(1, Ordering::Relaxed);
self.stats.remote_failures.fetch_add(1, Ordering::Relaxed);
}
}
#[derive(Debug, Clone)]
struct VerifiedBlob {
path: PathBuf,
len: u64,
modified: SystemTime,
}
impl VerifiedBlob {
fn describe(path: &Path) -> Option<Self> {
let metadata = std::fs::metadata(path).ok()?;
Some(Self {
path: path.to_path_buf(),
len: metadata.len(),
modified: metadata.modified().ok()?,
})
}
fn is_unchanged(&self) -> bool {
let Ok(metadata) = std::fs::metadata(&self.path) else {
return false;
};
metadata.len() == self.len && metadata.modified().is_ok_and(|now| now == self.modified)
}
}
#[derive(Debug, Clone, Default)]
struct TaskActionState {
manifest: String,
baseline_loaded: bool,
predictions: BTreeMap<CacheDigest, ActionPrediction>,
pending_predictions: BTreeMap<CacheDigest, ActionPrediction>,
remote_etag: Option<String>,
}
struct PrefetchedAction {
adapter: String,
result: RemoteActionResult,
}
struct RemoteDownloadReservation {
counter: Arc<AtomicU64>,
reserved: u64,
committed: bool,
}
impl RemoteDownloadReservation {
fn bytes(&self) -> u64 {
self.reserved
}
fn commit(mut self, bytes: u64) {
debug_assert!(bytes <= self.reserved);
self.counter
.fetch_sub(self.reserved.saturating_sub(bytes), Ordering::AcqRel);
self.committed = true;
}
}
impl Drop for RemoteDownloadReservation {
fn drop(&mut self) {
if !self.committed {
self.counter.fetch_sub(self.reserved, Ordering::AcqRel);
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
struct ExecutableIdentityKey {
executable: PathBuf,
environment: BTreeMap<String, Option<String>>,
}
impl CacheAgent {
pub fn new(cache_dir: impl Into<PathBuf>, version: impl Into<Arc<str>>) -> Self {
Self::build(cache_dir.into(), version.into(), None, 0)
}
pub fn new_remote(
cache_dir: impl Into<PathBuf>,
version: impl Into<Arc<str>>,
remote: AgentRemoteCache,
) -> Self {
Self::build(
cache_dir.into(),
version.into(),
Some(remote),
DEFAULT_MAX_REMOTE_DOWNLOAD_BYTES,
)
}
pub fn new_remote_with_download_limit(
cache_dir: impl Into<PathBuf>,
version: impl Into<Arc<str>>,
remote: AgentRemoteCache,
max_remote_download_bytes: u64,
) -> Self {
Self::build(
cache_dir.into(),
version.into(),
Some(remote),
max_remote_download_bytes,
)
}
fn build(
cache_dir: PathBuf,
version: Arc<str>,
remote: Option<AgentRemoteCache>,
remote_download_limit: u64,
) -> Self {
let remote_mode = remote
.as_ref()
.map_or(RemoteCacheMode::ReadOnly, |remote| remote.mode);
let remote_staging_dir = remote.as_ref().map_or_else(
|| cache_dir.join("remote"),
|remote| remote.staging_dir.clone(),
);
let remote = remote.map(|remote| Arc::new(remote.client));
let stats = Arc::new(AtomicAgentStats::default());
let remote_transfers = Arc::new(tokio::sync::Semaphore::new(MAX_REMOTE_TRANSFERS));
let uploads = remote
.clone()
.filter(|_| remote_mode.writes())
.map(|client| {
UploadQueue::new(
client,
Arc::new(AgentUploadSink {
stats: stats.clone(),
}),
remote_transfers.clone(),
)
});
Self {
cas: LocalCas::new(cache_dir.clone()),
actions: LocalActionCache::new(cache_dir.clone()),
verified_blobs: Arc::new(Mutex::new(BTreeMap::new())),
version,
write_locks: Arc::new(Mutex::new(BTreeMap::new())),
action_locks: Arc::new(Mutex::new(BTreeMap::new())),
stats,
observer: None,
executable_identities: Arc::new(Mutex::new(BTreeMap::new())),
manifest_dir: Arc::new(task_manifest_dir(&cache_dir)),
task_actions: Arc::new(Mutex::new(BTreeMap::new())),
next_task_run: Arc::new(AtomicU64::new(0)),
manifest_write_lock: Arc::new(Mutex::new(())),
remote,
remote_mode,
remote_staging_dir: Arc::new(remote_staging_dir),
remote_download_limit,
remote_download_bytes: Arc::new(AtomicU64::new(0)),
pending_remote_actions: Arc::new(Mutex::new(BTreeMap::new())),
remote_transfers,
prefetch_transfers: Arc::new(tokio::sync::Semaphore::new(MAX_PREFETCH_TRANSFERS)),
prefetch_tasks: Arc::new(Mutex::new(Vec::new())),
warnings: Arc::new(Mutex::new(BTreeSet::new())),
file_digests: Arc::new(Mutex::new(BTreeMap::new())),
uploads,
}
}
#[must_use]
pub fn with_observer(mut self, observer: Arc<dyn AgentEventObserver>) -> Self {
self.observer = Some(observer);
self
}
fn emit(&self, event: impl FnOnce() -> AgentEvent) {
if let Some(observer) = &self.observer {
observer.event(event());
}
}
fn reserve_remote_download(&self, bytes: u64) -> Result<RemoteDownloadReservation> {
let mut current = self.remote_download_bytes.load(Ordering::Acquire);
loop {
let next = current
.checked_add(bytes)
.ok_or_else(|| eyre::eyre!("remote cache download budget overflowed"))?;
if next > self.remote_download_limit {
bail!(
"remote cache download budget exceeded: {} bytes requested with {} of {} bytes already used",
bytes,
current,
self.remote_download_limit
);
}
match self.remote_download_bytes.compare_exchange_weak(
current,
next,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
return Ok(RemoteDownloadReservation {
counter: self.remote_download_bytes.clone(),
reserved: bytes,
committed: false,
});
}
Err(observed) => current = observed,
}
}
}
fn reserve_remote_download_up_to(&self, requested: u64) -> Result<RemoteDownloadReservation> {
let mut current = self.remote_download_bytes.load(Ordering::Acquire);
loop {
let reserved = requested.min(self.remote_download_limit.saturating_sub(current));
let next = current + reserved;
match self.remote_download_bytes.compare_exchange_weak(
current,
next,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
return Ok(RemoteDownloadReservation {
counter: self.remote_download_bytes.clone(),
reserved,
committed: false,
});
}
Err(observed) => current = observed,
}
}
}
pub async fn begin_task(&self, task: &str) -> Result<String> {
self.begin_task_with_remote_errors(task, false).await
}
pub async fn prefetch_task(&self, task: &str) -> Result<String> {
let run = self.begin_task_with_remote_errors(task, true).await?;
self.wait_for_prefetches().await;
Ok(run)
}
async fn begin_task_with_remote_errors(&self, task: &str, strict: bool) -> Result<String> {
validate_task_identity(task)?;
let early_predictions = {
let _write_guard = self.manifest_write_lock.lock().unwrap();
let _file_guard = self.lock_task_manifest(task)?;
self.load_task_manifest(task)?
.map(|manifest| manifest.predictions)
.unwrap_or_default()
};
let early_actions: BTreeSet<_> = early_predictions
.iter()
.map(|prediction| prediction.action.clone())
.collect();
self.spawn_prefetch_predictions(early_predictions);
let (remote_manifest, mut remote_etag) = if self.remote_mode.reads() {
match self.get_remote_task_manifest(task).await {
Ok(Some((manifest, etag))) => (Some(manifest), Some(etag)),
Ok(None) => (None, None),
Err(error) => {
if strict {
return Err(error).wrap_err_with(|| {
format!("remote task action manifest lookup failed for {task}")
});
}
self.note_remote_failure();
warn!("remote task action manifest lookup failed for {task}: {error}");
(None, None)
}
}
} else {
(None, None)
};
let manifest = {
let _write_guard = self.manifest_write_lock.lock().unwrap();
let _file_guard = self.lock_task_manifest(task)?;
let local_manifest = self.load_task_manifest(task)?;
let manifest = match (remote_manifest, local_manifest) {
(Some(remote), Some(local)) => {
let (manifest, merged) = merge_remote_task_manifest(task, remote, local);
if !merged {
remote_etag = None;
}
Some(manifest)
}
(Some(remote), None) => Some(remote),
(None, local) => local,
};
if let Some(manifest) = &manifest {
self.persist_task_manifest(manifest)?;
}
manifest
};
let state = if let Some(manifest) = manifest {
TaskActionState {
manifest: task.to_string(),
baseline_loaded: true,
predictions: manifest
.predictions
.into_iter()
.map(|prediction| (prediction.invocation.clone(), prediction))
.collect(),
pending_predictions: BTreeMap::new(),
remote_etag,
}
} else {
TaskActionState {
manifest: task.to_string(),
baseline_loaded: true,
remote_etag,
..TaskActionState::default()
}
};
let sequence = self.next_task_run.fetch_add(1, Ordering::Relaxed);
let run =
CacheDigest::blake3(format!("{task}\0{}\0{sequence}", std::process::id()).as_bytes())
.hash;
self.stats.predictions_loaded.fetch_max(
state.predictions.len().try_into().unwrap_or(u64::MAX),
Ordering::Relaxed,
);
let predictions = state
.predictions
.values()
.filter(|prediction| !early_actions.contains(&prediction.action))
.cloned()
.collect();
self.task_actions.lock().unwrap().insert(run.clone(), state);
self.spawn_prefetch_predictions(predictions);
Ok(run)
}
pub async fn cancel_prefetches(&self) {
let tasks = std::mem::take(&mut *self.prefetch_tasks.lock().unwrap());
for task in &tasks {
task.abort();
}
for task in tasks {
if let Err(error) = task.await
&& !error.is_cancelled()
{
warn!("remote action prefetch task failed: {error}");
}
}
}
pub async fn wait_for_uploads(&self) {
let Some(uploads) = &self.uploads else {
return;
};
let _timer = AtomicDurationTimer::start(&self.stats.upload_drain_duration_ns);
uploads.drain().await;
}
async fn wait_for_prefetches(&self) {
let tasks = std::mem::take(&mut *self.prefetch_tasks.lock().unwrap());
for task in tasks {
if let Err(error) = task.await {
warn!("remote action prefetch task failed: {error}");
}
}
}
pub async fn commit_task(&self, run: &str) -> Result<()> {
self.commit_task_actions(run).await.map(|_| ())
}
pub async fn commit_task_actions(&self, run: &str) -> Result<Vec<ActionPrediction>> {
validate_task_identity(run)?;
let state = self
.task_actions
.lock()
.unwrap()
.get(run)
.cloned()
.ok_or_else(|| eyre::eyre!("task action manifest baseline was not loaded"))?;
if !state.baseline_loaded {
bail!("task action manifest baseline was not loaded");
}
let task = state.manifest;
validate_task_identity(&task)?;
let completed = state
.pending_predictions
.values()
.cloned()
.collect::<Vec<_>>();
let (manifest, introduced) = {
let _write_guard = self.manifest_write_lock.lock().unwrap();
let _file_guard = self.lock_task_manifest(&task)?;
let mut predictions = self
.load_task_manifest(&task)?
.map(|manifest| {
manifest
.predictions
.into_iter()
.map(|prediction| (prediction.invocation.clone(), prediction))
.collect::<BTreeMap<_, _>>()
})
.unwrap_or_default();
let introduced: BTreeSet<CacheDigest> = state
.pending_predictions
.keys()
.filter(|invocation| !predictions.contains_key(*invocation))
.cloned()
.collect();
predictions.extend(state.pending_predictions);
let manifest = TaskActionManifest {
version: TASK_ACTION_MANIFEST_VERSION,
task: task.clone(),
predictions: predictions.into_values().collect(),
};
validate_task_manifest(&manifest, &task)?;
self.persist_task_manifest(&manifest)?;
(manifest, introduced)
};
self.task_actions.lock().unwrap().remove(run);
if self.remote_mode.writes() {
let mut manifest = manifest;
if let Some(uploads) = &self.uploads {
let actions: Vec<CacheDigest> = manifest
.predictions
.iter()
.map(|prediction| prediction.action.clone())
.collect();
let unpublished = uploads.wait_for_actions(&actions).await;
let withheld = manifest
.predictions
.iter()
.filter(|prediction| {
introduced.contains(&prediction.invocation)
&& unpublished.contains(&prediction.action)
})
.count();
if withheld > 0 {
warn!(
"{withheld} of {} predicted actions were not published, so the remote task action manifest omits them",
manifest.predictions.len()
);
manifest.predictions.retain(|prediction| {
!(introduced.contains(&prediction.invocation)
&& unpublished.contains(&prediction.action))
});
}
}
match self
.put_remote_task_manifest(&task, manifest, state.remote_etag)
.await
{
Ok(remote_manifest) => {
let _write_guard = self.manifest_write_lock.lock().unwrap();
let reconciliation = (|| {
let _file_guard = self.lock_task_manifest(&task)?;
let manifest = match self.load_task_manifest(&task)? {
Some(local) => {
merge_remote_task_manifest(&task, remote_manifest, local).0
}
None => remote_manifest,
};
self.persist_task_manifest(&manifest)
})();
if let Err(error) = reconciliation {
warn!(
"remote task action manifest reconciliation failed for {task}: {error}"
);
}
}
Err(error) => {
self.note_remote_failure();
warn!("remote task action manifest upload failed for {task}: {error}");
}
}
}
Ok(completed)
}
fn task_manifest_path(&self, task: &str) -> PathBuf {
self.manifest_dir.join(format!("{task}.json"))
}
fn task_manifest_lock_path(&self, task: &str) -> PathBuf {
self.manifest_dir.join("locks").join(format!("{task}.lock"))
}
fn lock_task_manifest(&self, task: &str) -> Result<fslock::LockFile> {
let path = self.task_manifest_lock_path(task);
fs::create_dir_all(path.parent().expect("task manifest lock has a parent"))?;
let mut lock = fslock::LockFile::open(&path)?;
lock.lock()?;
Ok(lock)
}
fn load_task_manifest(&self, task: &str) -> Result<Option<TaskActionManifest>> {
match fs::read(self.task_manifest_path(task)) {
Ok(contents) => Ok(Some(self.parse_task_manifest(task, &contents, false)?)),
Err(error) if error.kind() == std::io::ErrorKind::NotFound => Ok(None),
Err(error) => Err(error.into()),
}
}
fn parse_task_manifest(
&self,
task: &str,
contents: &[u8],
require_canonical: bool,
) -> Result<TaskActionManifest> {
let manifest: TaskActionManifest = serde_json::from_slice(contents)?;
validate_task_manifest(&manifest, task)?;
if require_canonical && canonical_json(&manifest)? != contents {
bail!("task action manifest is not canonical JSON");
}
Ok(manifest)
}
fn task_manifest_selector(task: &str) -> Result<(Vec<u8>, CacheDigest)> {
TaskActionManifest::selector(task)
}
fn persist_task_manifest(&self, manifest: &TaskActionManifest) -> Result<()> {
let bytes = canonical_json(manifest)?;
fs::create_dir_all(self.manifest_dir.as_path())?;
let mut temporary = tempfile::NamedTempFile::new_in(self.manifest_dir.as_path())?;
std::io::Write::write_all(temporary.as_file_mut(), &bytes)?;
temporary.as_file_mut().sync_all()?;
temporary
.persist(self.task_manifest_path(&manifest.task))
.map_err(|error| error.error)?;
Ok(())
}
async fn get_remote_task_manifest(
&self,
task: &str,
) -> Result<Option<(TaskActionManifest, String)>> {
let Some(remote) = &self.remote else {
return Ok(None);
};
let (_, selector) = Self::task_manifest_selector(task)?;
let _permit = self.remote_transfers.acquire().await?;
self.stats
.remote_manifest_lookups
.fetch_add(1, Ordering::Relaxed);
let _timer = AtomicDurationTimer::start(&self.stats.remote_manifest_lookup_duration_ns);
let Some(remote_manifest) = remote.get_action_manifest(&selector).await? else {
return Ok(None);
};
let manifest = self.parse_task_manifest(task, &remote_manifest.bytes, true)?;
Ok(Some((manifest, remote_manifest.etag)))
}
async fn put_remote_task_manifest(
&self,
task: &str,
mut manifest: TaskActionManifest,
mut expected_etag: Option<String>,
) -> Result<TaskActionManifest> {
let Some(remote) = &self.remote else {
return Ok(manifest);
};
let (_, selector) = Self::task_manifest_selector(task)?;
for _ in 0..4 {
let bytes = canonical_json(&manifest)?;
let outcome = {
let _permit = self.remote_transfers.acquire().await?;
remote
.put_action_manifest(&selector, &bytes, expected_etag.as_deref())
.await?
};
match outcome {
ManifestPutOutcome::Stored => return Ok(manifest),
ManifestPutOutcome::PreconditionFailed => {
let Some((remote_manifest, etag)) = self.get_remote_task_manifest(task).await?
else {
expected_etag = None;
continue;
};
manifest = merge_task_manifests(task, Some(remote_manifest), manifest)?;
expected_etag = Some(etag);
}
}
}
bail!("remote task action manifest changed too frequently")
}
pub fn stats(&self) -> AgentStats {
AgentStats {
session_duration_ns: 0,
lookups: self.stats.lookups.load(Ordering::Relaxed),
unconsulted: self.stats.unconsulted.load(Ordering::Relaxed),
hits: self.stats.hits.load(Ordering::Relaxed),
stores: self.stats.stores.load(Ordering::Relaxed),
stored_bytes: self.stats.stored_bytes.load(Ordering::Relaxed),
verifications: self.stats.verifications.load(Ordering::Relaxed),
divergences: self.stats.divergences.load(Ordering::Relaxed),
downloaded_bytes: self.stats.downloaded_bytes.load(Ordering::Relaxed),
uploaded_bytes: self.stats.uploaded_bytes.load(Ordering::Relaxed),
background_uploads: self.stats.background_uploads.load(Ordering::Relaxed),
background_upload_failures: self
.stats
.background_upload_failures
.load(Ordering::Relaxed),
remote_blob_pack_uploads: self.stats.remote_blob_pack_uploads.load(Ordering::Relaxed),
remote_blob_pack_upload_blobs: self
.stats
.remote_blob_pack_upload_blobs
.load(Ordering::Relaxed),
upload_drain_duration_ns: self.stats.upload_drain_duration_ns.load(Ordering::Relaxed),
prefetched_actions: self.stats.prefetched_actions.load(Ordering::Relaxed),
predictions_loaded: self.stats.predictions_loaded.load(Ordering::Relaxed),
bypasses: self.stats.bypasses.lock().unwrap().clone(),
avoided_compiler_duration_ns: self
.stats
.avoided_compiler_duration_ns
.load(Ordering::Relaxed),
compiler: self.stats.compiler.lock().unwrap().clone(),
slow_compilations: self.stats.slow_compilations.lock().unwrap().clone(),
remote_failures: self.stats.remote_failures.load(Ordering::Relaxed),
remote_manifest_lookups: self.stats.remote_manifest_lookups.load(Ordering::Relaxed),
remote_manifest_lookup_duration_ns: self
.stats
.remote_manifest_lookup_duration_ns
.load(Ordering::Relaxed),
remote_action_lookups: self.stats.remote_action_lookups.load(Ordering::Relaxed),
remote_action_lookup_duration_ns: self
.stats
.remote_action_lookup_duration_ns
.load(Ordering::Relaxed),
remote_blob_requests: self.stats.remote_blob_requests.load(Ordering::Relaxed),
remote_blob_pack_requests: self.stats.remote_blob_pack_requests.load(Ordering::Relaxed),
remote_blob_pack_blobs: self.stats.remote_blob_pack_blobs.load(Ordering::Relaxed),
remote_blob_transfer_duration_ns: self
.stats
.remote_blob_transfer_duration_ns
.load(Ordering::Relaxed),
local_cas_write_duration_ns: self
.stats
.local_cas_write_duration_ns
.load(Ordering::Relaxed),
prefetch_runs: self.stats.prefetch_runs.load(Ordering::Relaxed),
prefetch_duration_ns: self.stats.prefetch_duration_ns.load(Ordering::Relaxed),
materialization_duration_ns: self
.stats
.materialization_duration_ns
.load(Ordering::Relaxed),
restored_output_files: self.stats.restored_output_files.load(Ordering::Relaxed),
restored_output_bytes: self.stats.restored_output_bytes.load(Ordering::Relaxed),
reflinked_output_files: self.stats.reflinked_output_files.load(Ordering::Relaxed),
reflinked_output_bytes: self.stats.reflinked_output_bytes.load(Ordering::Relaxed),
copied_output_files: self.stats.copied_output_files.load(Ordering::Relaxed),
copied_output_bytes: self.stats.copied_output_bytes.load(Ordering::Relaxed),
reused_output_files: self.stats.reused_output_files.load(Ordering::Relaxed),
reused_output_bytes: self.stats.reused_output_bytes.load(Ordering::Relaxed),
}
}
pub async fn handle_requests(
&self,
requests: impl IntoIterator<Item = AgentRequest>,
) -> Vec<AgentResponse> {
let mut connection = ConnectionUploads::default();
let mut responses = Vec::new();
for request in requests {
responses.push(self.respond_on(request, &mut connection).await);
}
responses
}
fn write_lock(&self, digest: &CacheDigest) -> Arc<tokio::sync::Mutex<()>> {
Self::digest_lock(&self.write_locks, digest)
}
fn action_lock(&self, digest: &CacheDigest) -> Arc<tokio::sync::Mutex<()>> {
Self::digest_lock(&self.action_locks, digest)
}
fn digest_lock(
locks: &Mutex<BTreeMap<CacheDigest, Weak<tokio::sync::Mutex<()>>>>,
digest: &CacheDigest,
) -> Arc<tokio::sync::Mutex<()>> {
let mut locks = locks.lock().unwrap();
locks.retain(|_, lock| lock.strong_count() > 0);
if let Some(lock) = locks.get(digest).and_then(Weak::upgrade) {
return lock;
}
let lock = Arc::new(tokio::sync::Mutex::new(()));
locks.insert(digest.clone(), Arc::downgrade(&lock));
lock
}
#[cfg(test)]
async fn respond(&self, request: AgentRequest) -> AgentResponse {
self.respond_on(request, &mut ConnectionUploads::default())
.await
}
async fn respond_on(
&self,
request: AgentRequest,
connection: &mut ConnectionUploads,
) -> AgentResponse {
let result = match request {
AgentRequest::BeginTask { task } => self
.begin_task(&task)
.await
.map(|run| AgentResponse::TaskBegun { run }),
AgentRequest::CommitTask { run } => self
.commit_task(&run)
.await
.map(|()| AgentResponse::TaskCommitted),
AgentRequest::FindBlob { digest } => self.find_blob(&digest).await,
AgentRequest::FindBlobs { digests } => self.find_blobs(digests).await,
AgentRequest::StoreBlob { digest, source } => {
self.store_blob(&digest, &source, connection).await
}
AgentRequest::FindActionResult { action } => {
self.stats.lookups.fetch_add(1, Ordering::Relaxed);
self.find_action_result(&action).await
}
AgentRequest::RecordActionHit {
action,
restore,
crate_name,
} => self.record_action_hit(&action, restore, crate_name),
AgentRequest::RecordBypass { kind } => {
*self
.stats
.bypasses
.lock()
.unwrap()
.entry(kind.clone())
.or_insert(0) += 1;
self.emit(|| AgentEvent::Bypass { kind });
Ok(AgentResponse::BypassRecorded)
}
AgentRequest::RecordUnconsulted => {
self.stats.unconsulted.fetch_add(1, Ordering::Relaxed);
self.emit(|| AgentEvent::Unconsulted);
Ok(AgentResponse::UnconsultedRecorded)
}
AgentRequest::RecordWarning { message } => self.record_warning(message),
AgentRequest::FindFileDigests { scope, files } => self.find_file_digests(scope, files),
AgentRequest::JoinActionPromise {
adapter,
invocation,
} => self.join_action_promise(&adapter, &invocation).await,
AgentRequest::CompleteActionPromise { claim, prediction } => {
self.complete_action_promise(&claim, &prediction).await
}
AgentRequest::RecordFileDigests { scope, entries } => {
self.record_file_digests(scope, entries)
}
AgentRequest::RecordCompilerInvocation {
outcome,
crate_name,
duration_ns,
} => self.record_compiler_invocation(&outcome, crate_name.as_deref(), duration_ns),
AgentRequest::RecordActionVerification { matched, restore } => {
self.record_materialization(restore);
self.stats.verifications.fetch_add(1, Ordering::Relaxed);
if !matched {
self.stats.divergences.fetch_add(1, Ordering::Relaxed);
}
self.emit(|| AgentEvent::Verification { matched, restore });
Ok(AgentResponse::ActionVerificationRecorded)
}
AgentRequest::StoreActionResult { result } => {
self.store_action_result(&result, connection).await
}
AgentRequest::FindActionPrediction { task, invocation } => {
self.find_action_prediction(&task, &invocation)
}
AgentRequest::RecordActionPrediction { task, prediction } => {
self.record_action_prediction(&task, prediction)
}
AgentRequest::FindExecutableIdentity {
executable,
environment,
} => self.find_executable_identity(executable, environment),
AgentRequest::StoreExecutableIdentity {
executable,
environment,
stdout,
} => self.store_executable_identity(executable, environment, stdout),
AgentRequest::Hello { .. } => {
Err(eyre::eyre!("hello is only valid as the first request"))
}
};
result.unwrap_or_else(|error| AgentResponse::Error {
message: error.to_string(),
})
}
async fn find_blob(&self, digest: &CacheDigest) -> Result<AgentResponse> {
if let Some(path) = self.find_verified_blob(digest)? {
return Ok(AgentResponse::Blob { path: Some(path) });
}
if !self.remote_mode.reads() {
return Ok(AgentResponse::Blob { path: None });
}
let Some(remote) = &self.remote else {
return Ok(AgentResponse::Blob { path: None });
};
match self.fetch_remote_blob(remote, digest).await {
Ok(path) => Ok(AgentResponse::Blob { path: Some(path) }),
Err(error) => {
warn!(
"remote cache blob lookup failed for {}: {error}",
digest.hash
);
Ok(AgentResponse::Blob { path: None })
}
}
}
async fn find_blobs(&self, digests: Vec<CacheDigest>) -> Result<AgentResponse> {
let mut paths = BTreeMap::new();
let mut missing = Vec::new();
for digest in &digests {
match self.find_verified_blob(digest)? {
Some(path) => {
paths.insert(digest.clone(), path);
}
None => {
missing.push(digest.clone());
}
}
}
if !missing.is_empty()
&& self.remote_mode.reads()
&& let Some(remote) = &self.remote
{
paths.extend(self.fetch_remote_blobs(remote, missing, None).await);
}
Ok(AgentResponse::Blobs {
paths: digests
.into_iter()
.map(|digest| paths.get(&digest).cloned())
.collect(),
})
}
async fn store_blob(
&self,
digest: &CacheDigest,
source: &Path,
connection: &mut ConnectionUploads,
) -> Result<AgentResponse> {
let path = {
let lock = self.write_lock(digest);
let _guard = lock.lock().await;
if let Some(path) = self.find_verified_blob(digest)? {
path
} else {
let path = self.cas.store_file(digest, source)?;
self.remember_verified_blob(digest, &path);
self.stats.stores.fetch_add(1, Ordering::Relaxed);
self.stats
.stored_bytes
.fetch_add(digest.size, Ordering::Relaxed);
path
}
};
if let Some(uploads) = &self.uploads {
uploads.queue_blob(digest, path.clone(), connection);
}
Ok(AgentResponse::Stored { path })
}
fn find_verified_blob(&self, digest: &CacheDigest) -> Result<Option<PathBuf>> {
let remembered = self.verified_blobs.lock().unwrap().get(digest).cloned();
if let Some(remembered) = remembered {
if remembered.is_unchanged() {
return Ok(Some(remembered.path));
}
self.verified_blobs.lock().unwrap().remove(digest);
}
let path = self.cas.find(digest)?;
if let Some(path) = &path {
self.remember_verified_blob(digest, path);
}
Ok(path)
}
fn remember_verified_blob(&self, digest: &CacheDigest, path: &Path) {
let Some(verified) = VerifiedBlob::describe(path) else {
return;
};
self.verified_blobs
.lock()
.unwrap()
.insert(digest.clone(), verified);
}
async fn find_action_result(&self, action: &CacheDigest) -> Result<AgentResponse> {
if let Some(result) = self.actions.find(action)? {
return Ok(AgentResponse::ActionResult {
result: Some(result),
});
}
if !self.remote_mode.reads() {
return Ok(AgentResponse::ActionResult { result: None });
}
let Some(remote) = &self.remote else {
return Ok(AgentResponse::ActionResult { result: None });
};
let lock = self.action_lock(action);
let _guard = lock.lock().await;
if let Some(result) = self.actions.find(action)? {
return Ok(AgentResponse::ActionResult {
result: Some(result),
});
}
if let Some(result) = self
.pending_remote_actions
.lock()
.unwrap()
.get(action)
.cloned()
{
return Ok(AgentResponse::ActionResult {
result: Some(result),
});
}
let _permit = self.remote_transfers.acquire().await?;
match self.get_remote_action_result(remote, action).await {
Ok(Some(result)) => {
self.pending_remote_actions
.lock()
.unwrap()
.insert(action.clone(), result.clone());
Ok(AgentResponse::ActionResult {
result: Some(result),
})
}
Ok(None) => Ok(AgentResponse::ActionResult { result: None }),
Err(error) => {
self.note_remote_failure();
warn!(
"remote cache action lookup failed for {}: {error}",
action.hash
);
Ok(AgentResponse::ActionResult { result: None })
}
}
}
async fn store_action_result(
&self,
result: &RemoteActionResult,
connection: &ConnectionUploads,
) -> Result<AgentResponse> {
let path = self.actions.store(result)?;
if let Some(uploads) = &self.uploads {
uploads.queue_action_result(result, connection);
}
Ok(AgentResponse::ActionStored { path })
}
async fn join_action_promise(
&self,
adapter: &str,
invocation: &CacheDigest,
) -> Result<AgentResponse> {
if !self.remote_mode.reads() || !self.remote_mode.writes() {
return Ok(AgentResponse::ActionPromise {
claim: None,
prediction: None,
});
}
let Some(remote) = &self.remote else {
return Ok(AgentResponse::ActionPromise {
claim: None,
prediction: None,
});
};
let deadline = Instant::now() + ACTION_PROMISE_WAIT;
loop {
let _permit = self.remote_transfers.acquire().await?;
let state = remote.join_action_promise(invocation, adapter).await;
drop(_permit);
match state {
Ok(Some(ActionPromiseState::Claimed { claim })) => {
return Ok(AgentResponse::ActionPromise {
claim: Some(claim),
prediction: None,
});
}
Ok(Some(ActionPromiseState::Complete { prediction })) => {
return Ok(AgentResponse::ActionPromise {
claim: None,
prediction: Some(prediction),
});
}
Ok(Some(ActionPromiseState::Pending { retry_after_ms }))
if Instant::now() < deadline =>
{
tokio::time::sleep(Duration::from_millis(retry_after_ms.clamp(10, 5_000)))
.await;
}
Ok(Some(ActionPromiseState::Pending { .. }) | None) => {
return Ok(AgentResponse::ActionPromise {
claim: None,
prediction: None,
});
}
Ok(Some(_)) => {
return Ok(AgentResponse::ActionPromise {
claim: None,
prediction: None,
});
}
Err(error) => {
self.note_remote_failure();
warn!(
"remote cache action promise failed for {}: {error}",
invocation.hash
);
return Ok(AgentResponse::ActionPromise {
claim: None,
prediction: None,
});
}
}
}
}
async fn complete_action_promise(
&self,
claim: &str,
prediction: &ActionPrediction,
) -> Result<AgentResponse> {
prediction.validate()?;
let Some(remote) = &self.remote else {
return Ok(AgentResponse::ActionPromiseCompleted);
};
let Some(uploads) = &self.uploads else {
return Ok(AgentResponse::ActionPromiseCompleted);
};
if uploads
.wait_for_actions(std::slice::from_ref(&prediction.action))
.await
.contains(&prediction.action)
{
return Ok(AgentResponse::ActionPromiseCompleted);
}
let completion = ActionPromiseCompletion {
claim: claim.to_string(),
prediction: prediction.clone(),
};
let _permit = self.remote_transfers.acquire().await?;
if let Err(error) = remote
.complete_action_promise(&prediction.invocation, &completion)
.await
{
self.note_remote_failure();
warn!(
"remote cache action promise completion failed for {}: {error}",
prediction.invocation.hash
);
}
Ok(AgentResponse::ActionPromiseCompleted)
}
fn note_remote_failure(&self) {
self.stats.remote_failures.fetch_add(1, Ordering::Relaxed);
}
async fn get_remote_action_result(
&self,
remote: &RemoteCacheClient,
action: &CacheDigest,
) -> Result<Option<RemoteActionResult>> {
self.stats
.remote_action_lookups
.fetch_add(1, Ordering::Relaxed);
let _timer = AtomicDurationTimer::start(&self.stats.remote_action_lookup_duration_ns);
remote.get_action_result(action).await
}
fn record_action_hit(
&self,
action: &CacheDigest,
restore: RestoreStats,
crate_name: Option<String>,
) -> Result<AgentResponse> {
validate_crate_name(crate_name.as_deref())?;
if self.actions.find(action)?.is_none() {
let pending = self.pending_remote_actions.lock().unwrap().remove(action);
if let Some(result) = pending {
self.actions.store(&result)?;
} else {
bail!("cannot record a hit for a missing action result");
}
}
self.record_restore(restore);
self.stats.hits.fetch_add(1, Ordering::Relaxed);
self.emit(|| AgentEvent::ActionHit {
crate_name,
restore,
});
Ok(AgentResponse::ActionHitRecorded)
}
fn record_restore(&self, restore: RestoreStats) {
self.record_materialization(restore);
atomic_saturating_add(
&self.stats.avoided_compiler_duration_ns,
restore.avoided_compiler_duration_ns,
);
atomic_saturating_add(&self.stats.restored_output_files, restore.output_files);
atomic_saturating_add(&self.stats.restored_output_bytes, restore.output_bytes);
atomic_saturating_add(
&self.stats.reflinked_output_files,
restore.reflinked_output_files,
);
atomic_saturating_add(
&self.stats.reflinked_output_bytes,
restore.reflinked_output_bytes,
);
atomic_saturating_add(&self.stats.copied_output_files, restore.copied_output_files);
atomic_saturating_add(&self.stats.copied_output_bytes, restore.copied_output_bytes);
atomic_saturating_add(&self.stats.reused_output_files, restore.reused_output_files);
atomic_saturating_add(&self.stats.reused_output_bytes, restore.reused_output_bytes);
}
fn record_compiler_invocation(
&self,
outcome: &str,
crate_name: Option<&str>,
duration_ns: u64,
) -> Result<AgentResponse> {
if !matches!(
outcome,
"miss" | "unconsulted" | "bypass" | "verification" | "incremental"
) {
bail!("invalid compiler invocation outcome");
}
validate_crate_name(crate_name)?;
let mut compiler = self.stats.compiler.lock().unwrap();
let stats = compiler.entry(outcome.to_string()).or_default();
stats.invocations = stats.invocations.saturating_add(1);
stats.duration_ns = stats.duration_ns.saturating_add(duration_ns);
drop(compiler);
if outcome != "verification"
&& let Some(crate_name) = crate_name.filter(|name| !name.is_empty())
{
let mut slow = self.stats.slow_compilations.lock().unwrap();
let duration = slow.entry(crate_name.to_string()).or_default();
*duration = duration.saturating_add(duration_ns);
}
self.emit(|| AgentEvent::CompilerInvocation {
outcome: outcome.to_string(),
crate_name: crate_name.map(str::to_string),
duration_ns,
});
Ok(AgentResponse::CompilerInvocationRecorded)
}
fn record_materialization(&self, restore: RestoreStats) {
atomic_saturating_add(&self.stats.materialization_duration_ns, restore.duration_ns);
}
fn record_warning(&self, message: String) -> Result<AgentResponse> {
if message.is_empty()
|| message.len() > MAX_WARNING_BYTES
|| message.contains(['\n', '\r', '\0'])
{
bail!("invalid shim warning");
}
let mut warnings = self.warnings.lock().unwrap();
if !warnings.contains(&message) && warnings.len() < MAX_WARNINGS {
eprintln!("mbx[warning]: {message}");
self.emit(|| AgentEvent::Warning {
message: message.clone(),
});
warnings.insert(message);
}
Ok(AgentResponse::WarningRecorded)
}
fn find_file_digests(
&self,
scope: FileDigestScope,
files: Vec<FileIdentity>,
) -> Result<AgentResponse> {
if files.len() > MAX_FILE_DIGEST_BATCH {
bail!("too many file-digest lookups in one request");
}
let ledger = self.file_digests.lock().unwrap();
let digests = files
.into_iter()
.map(|file| {
let recorded = ledger.get(&(scope, file.path.clone()))?;
(recorded.file == file).then(|| recorded.digest.clone())
})
.collect();
Ok(AgentResponse::FileDigests { digests })
}
fn record_file_digests(
&self,
scope: FileDigestScope,
entries: Vec<RecordedFileDigest>,
) -> Result<AgentResponse> {
if entries.len() > MAX_FILE_DIGEST_BATCH {
bail!("too many file-digest records in one request");
}
for entry in &entries {
if !entry.file.path.is_absolute() {
bail!("file-digest records need absolute paths");
}
entry.digest.validate()?;
if entry.file.len != entry.digest.size {
bail!("file-digest record length does not match its digest");
}
}
let mut ledger = self.file_digests.lock().unwrap();
for entry in entries {
if ledger.len() >= MAX_FILE_DIGEST_ENTRIES
&& !ledger.contains_key(&(scope, entry.file.path.clone()))
{
break;
}
ledger.insert((scope, entry.file.path.clone()), entry);
}
Ok(AgentResponse::FileDigestsRecorded)
}
fn find_action_prediction(
&self,
task: &str,
invocation: &CacheDigest,
) -> Result<AgentResponse> {
validate_task_identity(task)?;
invocation.validate()?;
let prediction = self
.task_actions
.lock()
.unwrap()
.get(task)
.and_then(|state| state.predictions.get(invocation))
.cloned();
Ok(AgentResponse::ActionPrediction { prediction })
}
fn record_action_prediction(
&self,
task: &str,
prediction: ActionPrediction,
) -> Result<AgentResponse> {
validate_task_identity(task)?;
prediction.validate()?;
let mut tasks = self.task_actions.lock().unwrap();
let state = tasks.entry(task.to_string()).or_default();
if !state.predictions.contains_key(&prediction.invocation)
&& state.predictions.len() >= MAX_TASK_ACTION_PREDICTIONS
{
bail!("task action manifest contains too many predictions");
}
state
.predictions
.insert(prediction.invocation.clone(), prediction.clone());
state
.pending_predictions
.insert(prediction.invocation.clone(), prediction);
Ok(AgentResponse::ActionPredictionRecorded)
}
fn executable_identity_key(
&self,
executable: PathBuf,
environment: BTreeMap<String, Option<String>>,
) -> Result<ExecutableIdentityKey> {
if !environment.keys().all(|name| {
matches!(
name.as_str(),
"RUSTUP_HOME"
| "RUSTUP_TOOLCHAIN"
| "SDKROOT"
| "MACOSX_DEPLOYMENT_TARGET"
| "LIB"
| "UCRTVersion"
| "UniversalCRTSdkDir"
| "VCToolsInstallDir"
| "VCToolsVersion"
| "WindowsSdkDir"
| "WindowsSDKVersion"
)
}) {
bail!("executable identity contains an unsupported environment variable");
}
Ok(ExecutableIdentityKey {
executable,
environment,
})
}
fn find_executable_identity(
&self,
executable: PathBuf,
environment: BTreeMap<String, Option<String>>,
) -> Result<AgentResponse> {
let key = self.executable_identity_key(executable, environment)?;
let stdout = self
.executable_identities
.lock()
.unwrap()
.get(&key)
.cloned();
Ok(AgentResponse::ExecutableIdentity { stdout })
}
fn store_executable_identity(
&self,
executable: PathBuf,
environment: BTreeMap<String, Option<String>>,
stdout: Vec<u8>,
) -> Result<AgentResponse> {
if stdout.len() > MAX_EXECUTABLE_IDENTITY_SIZE {
bail!("executable identity exceeds {MAX_EXECUTABLE_IDENTITY_SIZE} bytes");
}
let key = self.executable_identity_key(executable, environment)?;
let mut identities = self.executable_identities.lock().unwrap();
let is_new = !identities.contains_key(&key);
let previous_size = identities.get(&key).map_or(0, Vec::len);
if is_new && identities.len() >= MAX_EXECUTABLE_IDENTITIES {
bail!("executable identity cache contains too many entries");
}
let retained_bytes = identities.values().map(Vec::len).sum::<usize>();
if retained_bytes - previous_size + stdout.len() > MAX_EXECUTABLE_IDENTITY_BYTES {
bail!("executable identity cache contains too many bytes");
}
identities.insert(key, stdout.clone());
Ok(AgentResponse::ExecutableIdentity {
stdout: Some(stdout),
})
}
pub async fn handle_connection<S>(&self, stream: S) -> Result<()>
where
S: AsyncRead + AsyncWrite + Unpin,
{
let (reader, mut writer) = tokio::io::split(stream);
let mut reader = BufReader::new(reader);
let hello = read_request(&mut reader)
.await?
.ok_or_else(|| eyre::eyre!("connection closed before the agent handshake"))?;
let request: AgentRequest = serde_json::from_str(&hello)?;
match request {
AgentRequest::Hello {
protocol,
client_version,
} if protocol == AGENT_PROTOCOL_VERSION && client_version == self.version.as_ref() => {}
AgentRequest::Hello { protocol, .. } if protocol != AGENT_PROTOCOL_VERSION => {
send_response(
&mut writer,
&AgentResponse::Error {
message: format!(
"unsupported agent protocol {protocol}; expected {AGENT_PROTOCOL_VERSION}"
),
},
)
.await?;
return Ok(());
}
AgentRequest::Hello { client_version, .. } => {
send_response(
&mut writer,
&AgentResponse::Error {
message: format!(
"cache client {client_version} does not match agent {}",
self.version
),
},
)
.await?;
return Ok(());
}
_ => bail!("the first agent request must be hello"),
}
send_response(
&mut writer,
&AgentResponse::Hello {
protocol: AGENT_PROTOCOL_VERSION,
agent_version: self.version.to_string(),
},
)
.await?;
let mut connection = ConnectionUploads::default();
while let Some(line) = read_request(&mut reader).await? {
let response = match serde_json::from_str(&line) {
Ok(request) => self.respond_on(request, &mut connection).await,
Err(error) => AgentResponse::Error {
message: format!("invalid agent request: {error}"),
},
};
send_response(&mut writer, &response).await?;
}
Ok(())
}
}
async fn read_request<R>(reader: &mut R) -> Result<Option<String>>
where
R: AsyncBufRead + Unpin,
{
let mut line = Vec::new();
loop {
let available = reader.fill_buf().await?;
if available.is_empty() {
break;
}
let (consumed, complete) = match available.iter().position(|byte| *byte == b'\n') {
Some(index) => (index, true),
None => (available.len(), false),
};
if line.len() + consumed > MAX_REQUEST_BYTES {
bail!("agent request exceeded {MAX_REQUEST_BYTES} bytes");
}
line.extend_from_slice(&available[..consumed]);
reader.consume(consumed + usize::from(complete));
if complete {
return Ok(Some(String::from_utf8(line)?));
}
}
if line.is_empty() {
Ok(None)
} else {
Ok(Some(String::from_utf8(line)?))
}
}
async fn send_response(
writer: &mut (impl AsyncWrite + Unpin),
response: &AgentResponse,
) -> Result<()> {
let mut encoded = serde_json::to_vec(response)?;
encoded.push(b'\n');
writer.write_all(&encoded).await?;
writer.flush().await?;
Ok(())
}
#[cfg(test)]
#[path = "agent_tests.rs"]
mod tests;