use crate::id::KernelId;
use crate::memory_management::FailureId;
use crate::server::{BufferBinding, ServerError};
use crate::stream::{FailureStore, StreamCapture};
use alloc::vec::Vec;
use cubecl_environment::stream::StreamId;
pub trait WriteScoped: Sized {
type Streams: FailureStore;
fn write_streams(&mut self) -> &mut Self::Streams;
#[allow(unused_variables)]
fn on_failure(&mut self, stream: StreamId, error: &ServerError) {}
#[allow(unused_variables)]
fn capturing(&mut self, stream: StreamId) -> Option<&mut StreamCapture> {
None
}
fn write_set(&mut self) -> Vec<BufferBinding> {
self.write_streams().write_set()
}
}
#[derive(Debug)]
pub enum ScopedOutcome<R> {
Executed(R),
Skipped,
Failed(ServerError),
}
impl<R> ScopedOutcome<R> {
pub fn into_result(self) -> Result<R, ServerError> {
match self {
ScopedOutcome::Executed(result) => Ok(result),
ScopedOutcome::Failed(error) => Err(error),
ScopedOutcome::Skipped => Err(ServerError::Skipped),
}
}
}
enum Opened {
Entered {
provisional: Option<FailureId>,
written: Vec<BufferBinding>,
},
Skipped,
}
pub struct ExecuteScope<'a, S: WriteScoped> {
server: &'a mut S,
stream: StreamId,
opened: Opened,
}
impl<'a, S: WriteScoped> ExecuteScope<'a, S> {
pub fn over(server: &'a mut S, stream: StreamId, written: Vec<BufferBinding>) -> Self {
let provisional = server.write_streams().enter_write(&written);
Self {
server,
stream,
opened: Opened::Entered {
provisional,
written,
},
}
}
pub fn launching<'b>(
server: &'a mut S,
kernel: KernelId,
stream: StreamId,
reads: impl Iterator<Item = &'b BufferBinding>,
written: Vec<BufferBinding>,
) -> Self {
let Some(found) = server.write_streams().read_failure(reads) else {
return Self::over(server, stream, written);
};
server.on_failure(stream, &found.error);
if let Some(capture) = server.capturing(stream) {
capture.fail(found.error.clone());
}
server.write_streams().propagate(&found, kernel, written);
Self {
server,
stream,
opened: Opened::Skipped,
}
}
pub fn skipped(&self) -> bool {
matches!(self.opened, Opened::Skipped)
}
pub fn execute<R>(
self,
body: impl FnOnce(&mut S) -> Result<R, ServerError>,
) -> ScopedOutcome<R> {
let Opened::Entered {
provisional,
written,
} = self.opened
else {
return ScopedOutcome::Skipped;
};
let result = body(self.server);
if let Some(capture) = self.server.capturing(self.stream) {
match result.as_ref() {
Ok(_) => capture.record(written.iter().cloned()),
Err(error) => capture.fail(error.clone()),
}
}
self.server
.write_streams()
.exit_write(provisional, written, result.as_ref().err());
match result {
Ok(result) => ScopedOutcome::Executed(result),
Err(error) => {
self.server.on_failure(self.stream, &error);
ScopedOutcome::Failed(error)
}
}
}
}
pub fn failed_writing<S: WriteScoped>(
server: &mut S,
stream: StreamId,
written: Vec<BufferBinding>,
error: ServerError,
) {
let _ = ExecuteScope::over(server, stream, written).execute(|_| Err::<(), _>(error));
}