use super::{
DeviceLossReport, FailureReason, ResidencyBudget, ResidencyError, ResidencyKey,
ResidencyMachine, ResidencyOutput, ResidencyPhase, ResidencyRequest, ResidencyTicket,
StaleCompletion,
};
use crate::{ChunkData, ChunkFootprint, ChunkId, DatasetId};
use std::collections::BTreeMap;
use thiserror::Error;
#[cfg(test)]
#[path = "host_working_set_tests.rs"]
mod tests;
#[cfg(test)]
#[path = "host_working_set_payload_tests.rs"]
mod payload_tests;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Error)]
pub enum HostWorkingSetError {
#[error("delivered chunk belongs to dataset {actual}, expected {expected}")]
DatasetMismatch {
expected: DatasetId,
actual: DatasetId,
},
#[error("delivered chunk is {actual}, expected {expected}")]
ChunkMismatch {
expected: ChunkId,
actual: ChunkId,
},
#[error("delivered chunk footprint {actual:?} differs from requested {expected:?}")]
FootprintMismatch {
expected: ChunkFootprint,
actual: ChunkFootprint,
},
#[error("stale residency completion for generation {}", .0.ticket.generation())]
StaleCompletion(StaleCompletion),
#[error(transparent)]
Residency(#[from] ResidencyError),
}
#[derive(Debug)]
pub struct HostWorkingSet {
machine: ResidencyMachine,
payloads: BTreeMap<ResidencyKey, ChunkData>,
}
impl HostWorkingSet {
#[must_use]
pub fn new(budget: ResidencyBudget) -> Self {
Self {
machine: ResidencyMachine::new(budget),
payloads: BTreeMap::new(),
}
}
#[must_use]
pub const fn residency(&self) -> &ResidencyMachine {
&self.machine
}
#[must_use]
pub fn retained_payloads(&self) -> usize {
self.payloads.len()
}
#[must_use]
pub fn payload(&self, key: ResidencyKey) -> Option<&ChunkData> {
self.payloads.get(&key)
}
pub fn request_into(
&mut self,
request: ResidencyRequest,
output: &mut ResidencyOutput,
) -> Result<ResidencyTicket, HostWorkingSetError> {
let ticket = self.machine.request_into(request, output)?;
self.payloads.remove(&request.key);
self.reconcile(output);
Ok(ticket)
}
pub fn deliver_into(
&mut self,
ticket: ResidencyTicket,
payload: ChunkData,
output: &mut ResidencyOutput,
) -> Result<(), HostWorkingSetError> {
validate_identity(ticket.key, &payload)?;
self.validate_footprint(ticket, &payload)?;
let transition = self.machine.ready_cpu_into(ticket, output);
self.reconcile(output);
let accepted = transition?;
if !accepted {
return Err(stale_error(ticket, &self.machine, output));
}
self.payloads.insert(ticket.key, payload);
self.reconcile(output);
Ok(())
}
pub fn begin_upload_into(
&mut self,
ticket: ResidencyTicket,
output: &mut ResidencyOutput,
) -> Result<(), HostWorkingSetError> {
let transition = self.machine.begin_upload_into(ticket, output);
self.reconcile(output);
let accepted = transition?;
self.require_current(ticket, accepted, output)
}
pub fn defer_upload_into(
&mut self,
ticket: ResidencyTicket,
output: &mut ResidencyOutput,
) -> Result<(), HostWorkingSetError> {
let transition = self.machine.defer_upload_into(ticket, output);
self.reconcile(output);
let accepted = transition?;
self.require_current(ticket, accepted, output)
}
pub fn complete_upload_into(
&mut self,
ticket: ResidencyTicket,
output: &mut ResidencyOutput,
) -> Result<(), HostWorkingSetError> {
let transition = self.machine.complete_upload_into(ticket, output);
self.reconcile(output);
let accepted = transition?;
self.require_current(ticket, accepted, output)
}
pub fn cancel_into(
&mut self,
ticket: ResidencyTicket,
output: &mut ResidencyOutput,
) -> Result<(), HostWorkingSetError> {
let accepted = self.machine.cancel_into(ticket, output);
self.require_current(ticket, accepted, output)
}
pub fn fail_into(
&mut self,
ticket: ResidencyTicket,
reason: FailureReason,
output: &mut ResidencyOutput,
) -> Result<(), HostWorkingSetError> {
let accepted = self.machine.fail_into(ticket, reason, output);
self.require_current(ticket, accepted, output)
}
pub fn set_budget_into(&mut self, budget: ResidencyBudget, output: &mut ResidencyOutput) {
self.machine.set_budget_into(budget, output);
self.reconcile(output);
self.payloads
.retain(|key, _| self.machine.snapshot(*key).phase != ResidencyPhase::Absent);
}
pub fn device_lost_into(
&mut self,
output: &mut ResidencyOutput,
) -> Result<DeviceLossReport, HostWorkingSetError> {
let report = self.machine.device_lost_into(output)?;
self.reconcile(output);
Ok(report)
}
fn require_current(
&mut self,
ticket: ResidencyTicket,
accepted: bool,
output: &ResidencyOutput,
) -> Result<(), HostWorkingSetError> {
self.drop_if_absent(ticket.key);
self.reconcile(output);
if accepted {
Ok(())
} else {
Err(stale_error(ticket, &self.machine, output))
}
}
fn validate_footprint(
&self,
ticket: ResidencyTicket,
payload: &ChunkData,
) -> Result<(), HostWorkingSetError> {
let snapshot = self.machine.snapshot(ticket.key);
if snapshot.generation != ticket.generation() {
return Ok(());
}
let Some(expected) = self.machine.footprint(ticket.key) else {
return Ok(());
};
let actual = payload.footprint();
if expected != actual {
return Err(HostWorkingSetError::FootprintMismatch { expected, actual });
}
Ok(())
}
fn reconcile(&mut self, output: &ResidencyOutput) {
for cancellation in &output.cancellations {
self.drop_if_absent(cancellation.key);
}
for eviction in &output.evictions {
self.drop_if_absent(eviction.key);
}
}
fn drop_if_absent(&mut self, key: ResidencyKey) {
if self.machine.snapshot(key).phase == ResidencyPhase::Absent {
self.payloads.remove(&key);
}
}
}
fn validate_identity(key: ResidencyKey, payload: &ChunkData) -> Result<(), HostWorkingSetError> {
if payload.dataset_id() != key.dataset {
return Err(HostWorkingSetError::DatasetMismatch {
expected: key.dataset,
actual: payload.dataset_id(),
});
}
if payload.chunk_id() != key.chunk {
return Err(HostWorkingSetError::ChunkMismatch {
expected: key.chunk,
actual: payload.chunk_id(),
});
}
Ok(())
}
fn stale_error(
ticket: ResidencyTicket,
machine: &ResidencyMachine,
output: &ResidencyOutput,
) -> HostWorkingSetError {
let stale = if let Some(stale) = output.stale.first().copied() {
stale
} else {
let current = machine.snapshot(ticket.key);
StaleCompletion {
ticket,
current_generation: current.generation,
current_phase: current.phase,
}
};
HostWorkingSetError::StaleCompletion(stale)
}