use crate::id::KernelId;
use crate::logging::ServerLogger;
use crate::memory_management::{Claim, ErrorGraph, FailureId, Skipped};
use crate::server::{BufferBinding, ServerError};
use crate::stream::{ReadFailure, StreamFactory, StreamMemory, StreamPool, base};
use alloc::sync::Arc;
use alloc::vec::Vec;
#[derive(Debug)]
pub struct Failures {
graph: ErrorGraph,
scratch: Vec<BufferBinding>,
logger: Arc<ServerLogger>,
}
impl Failures {
pub fn new(logger: Arc<ServerLogger>) -> Self {
Self {
graph: ErrorGraph::default(),
scratch: Vec::new(),
logger,
}
}
pub fn graph(&self) -> &ErrorGraph {
&self.graph
}
pub fn graph_mut(&mut self) -> &mut ErrorGraph {
&mut self.graph
}
}
pub trait FailureStore {
type Factory: StreamFactory<Stream: StreamMemory>;
fn split(&mut self) -> (&mut StreamPool<Self::Factory>, &mut Failures);
fn parts(&self) -> (&StreamPool<Self::Factory>, &Failures);
fn ensure_written<'a>(
&self,
handles: impl Iterator<Item = &'a BufferBinding>,
) -> Result<(), ServerError> {
let (pool, failures) = self.parts();
failures.graph.reports(handles.filter_map(|handle| {
let memory = handle.memory.id();
if !handle.memory.descriptor().is_allocated() {
return Some(Claim::Unallocated(memory));
}
let failure = pool.try_get(&handle.stream)?.failure(handle)?;
Some(Claim::Failed(failure, memory))
}))
}
fn read_failure<'a>(
&self,
mut reads: impl Iterator<Item = &'a BufferBinding>,
) -> Option<ReadFailure> {
let (pool, failures) = self.parts();
reads.find_map(|handle| {
let failure = pool.try_get(&handle.stream)?.failure(handle)?;
Some(ReadFailure {
failure,
needed: handle.memory.id(),
error: failures.graph.error(failure)?.clone(),
})
})
}
fn taint<'a>(&mut self, error: ServerError, written: impl Iterator<Item = &'a BufferBinding>) {
let (pool, failures) = self.split();
base::taint(pool, error, written, &mut failures.graph);
}
fn written<'a>(&mut self, written: impl Iterator<Item = &'a BufferBinding>) {
let (pool, failures) = self.split();
base::written(pool, written, &mut failures.graph);
}
fn propagate(
&mut self,
found: &ReadFailure,
kernel: KernelId,
mut written: Vec<BufferBinding>,
) {
let (pool, failures) = self.split();
failures.graph.skipped(
found.failure,
Skipped {
kernel,
needed: found.needed,
produced: written.iter().map(|handle| handle.memory.id()).collect(),
},
);
base::taint_with(pool, found.failure, written.iter(), &mut failures.graph);
written.clear();
failures.scratch = written;
}
fn write_set(&mut self) -> Vec<BufferBinding> {
let (_, failures) = self.split();
core::mem::take(&mut failures.scratch)
}
fn enter_write(&mut self, written: &[BufferBinding]) -> Option<FailureId> {
if written.is_empty() {
return None;
}
let (pool, failures) = self.split();
let provisional = failures.graph.insert(ServerError::TornDown);
base::taint_with(pool, provisional, written.iter(), &mut failures.graph);
Some(provisional)
}
fn exit_write(
&mut self,
provisional: Option<FailureId>,
mut written: Vec<BufferBinding>,
error: Option<&ServerError>,
) {
match error {
None => self.written(written.iter()),
Some(error) => {
let (_, failures) = self.split();
failures.logger.log_failure(error);
if let Some(provisional) = provisional {
failures.graph.replace(provisional, error.clone());
}
}
}
let (_, failures) = self.split();
if let Some(provisional) = provisional {
failures.graph.prune(provisional);
}
written.clear();
failures.scratch = written;
}
}