use eredu_core::{execution_control::*, generation::FinishReason};
use std::{cell::RefCell, rc::Rc};
mod choice;
mod sampling;
mod snapshot;
pub use choice::{TokenChoiceController, TokenChoiceError};
pub use sampling::{
apply_prepared_sampling_override, apply_sampling_override, SamplingOverride,
SamplingOverrideError, SamplingStateFacts, TextSamplingControlBackend,
ValidatedSamplingOverride,
};
pub fn intervention_plan_storage_bytes(
plan: &eredu_core::intervention::InterventionPlan,
) -> Option<u64> {
(std::mem::size_of_val(plan) as u64).checked_add(storage::heap_bytes(plan)?)
}
pub(crate) mod storage;
pub use snapshot::{
ManagedTextContinuation, SnapshotTokenController, TextBranchRequest, TextContinuationBranch,
TextContinuationSnapshot, TextSnapshotBackend, TextSnapshotError,
};
#[derive(Debug, Clone)]
pub struct GenerationBoundary {
status: GenerationStatus,
prediction: u64,
finish_reason: Option<FinishReason>,
}
#[derive(Debug)]
pub struct GenerationLifecycle {
boundary: GenerationBoundary,
epoch: u64,
}
impl Default for GenerationLifecycle {
fn default() -> Self {
Self {
boundary: GenerationBoundary {
status: GenerationStatus::Prepared,
prediction: 0,
finish_reason: None,
},
epoch: 0,
}
}
}
impl GenerationLifecycle {
pub fn status(&self) -> GenerationStatus {
self.boundary.status
}
pub fn next_prediction(&self) -> u64 {
self.boundary.prediction
}
pub fn epoch(&self) -> u64 {
self.epoch
}
pub fn finish_reason(&self) -> Option<FinishReason> {
self.boundary.finish_reason
}
fn invalid(&self, to: GenerationStatus) -> ExecutionControlError {
ExecutionControlError::Transition {
from: self.status(),
to,
}
}
pub fn begin_prediction(&mut self) -> Result<(), ExecutionControlError> {
if !matches!(
self.status(),
GenerationStatus::Prepared | GenerationStatus::Paused
) {
return Err(self.invalid(GenerationStatus::Running));
}
self.boundary
.prediction
.checked_add(1)
.ok_or(ExecutionControlError::Overflow)?;
self.boundary.status = GenerationStatus::Running;
Ok(())
}
pub fn complete_prediction(
&mut self,
reason: Option<FinishReason>,
) -> Result<(), ExecutionControlError> {
if self.status() != GenerationStatus::Running {
return Err(self.invalid(GenerationStatus::Paused));
}
self.boundary.prediction = self
.boundary
.prediction
.checked_add(1)
.ok_or(ExecutionControlError::Overflow)?;
self.boundary.finish_reason = reason;
self.boundary.status = match reason {
Some(FinishReason::Cancelled) => GenerationStatus::Cancelled,
Some(_) => GenerationStatus::Completed,
None => GenerationStatus::Paused,
};
Ok(())
}
pub fn pause(&mut self) -> Result<(), ExecutionControlError> {
if !matches!(
self.status(),
GenerationStatus::Prepared | GenerationStatus::Paused
) {
return Err(self.invalid(GenerationStatus::Paused));
}
self.boundary.status = GenerationStatus::Paused;
Ok(())
}
pub fn cancel(&mut self) -> Result<(), ExecutionControlError> {
if !matches!(
self.status(),
GenerationStatus::Prepared | GenerationStatus::Paused
) {
return Err(self.invalid(GenerationStatus::Cancelled));
}
self.boundary.status = GenerationStatus::Cancelled;
self.boundary.finish_reason = Some(FinishReason::Cancelled);
Ok(())
}
pub fn cancel_without_prediction(&mut self) -> Result<(), ExecutionControlError> {
if self.status() != GenerationStatus::Running {
return Err(self.invalid(GenerationStatus::Cancelled));
}
self.boundary.status = GenerationStatus::Cancelled;
self.boundary.finish_reason = Some(FinishReason::Cancelled);
Ok(())
}
pub fn fail(&mut self) {
self.boundary.status = GenerationStatus::Failed;
}
pub fn checkpoint(&self) -> Result<GenerationBoundary, ExecutionControlError> {
if !matches!(
self.status(),
GenerationStatus::Prepared | GenerationStatus::Paused | GenerationStatus::Completed
) {
return Err(self.invalid(GenerationStatus::Paused));
}
Ok(self.boundary.clone())
}
pub fn validate_restore(&self) -> Result<(), ExecutionControlError> {
self.checkpoint()?;
self.epoch
.checked_add(1)
.ok_or(ExecutionControlError::Overflow)?;
Ok(())
}
pub fn restore(&mut self, saved: &GenerationBoundary) -> Result<(), ExecutionControlError> {
self.validate_restore()?;
self.epoch += 1;
self.boundary = saved.clone();
Ok(())
}
pub fn fork(saved: &GenerationBoundary) -> Self {
Self {
boundary: saved.clone(),
epoch: 0,
}
}
}
struct BudgetState {
limits: SnapshotLimits,
usage: SnapshotUsage,
}
#[derive(Clone)]
pub struct SnapshotBudget(Rc<RefCell<BudgetState>>);
impl SnapshotBudget {
pub fn new(limits: SnapshotLimits) -> Self {
Self(Rc::new(RefCell::new(BudgetState {
limits,
usage: SnapshotUsage::default(),
})))
}
pub fn usage(&self) -> SnapshotUsage {
self.0.borrow().usage
}
pub fn reserve(
&self,
kind: SnapshotResourceKind,
estimate: Option<SnapshotEstimate>,
) -> Result<SnapshotReservation, ExecutionControlError> {
let estimate = estimate.ok_or(ExecutionControlError::UnknownEstimate)?;
let mut state = self.0.borrow_mut();
let add = |a: u64, b: u64| a.checked_add(b).ok_or(ExecutionControlError::Overflow);
let next = SnapshotUsage {
snapshots: add(
state.usage.snapshots,
u64::from(kind == SnapshotResourceKind::Snapshot),
)?,
branches: add(
state.usage.branches,
u64::from(kind == SnapshotResourceKind::Branch),
)?,
retained_bytes: add(state.usage.retained_bytes, estimate.retained_bytes)?,
cumulative_copy_bytes: add(state.usage.cumulative_copy_bytes, estimate.copy_bytes)?,
};
for (exceeded, name) in [
(
next.snapshots > state.limits.max_snapshots,
"snapshot count",
),
(next.branches > state.limits.max_branches, "branch count"),
(
next.retained_bytes > state.limits.retained_bytes,
"retained bytes",
),
(
next.cumulative_copy_bytes > state.limits.cumulative_copy_bytes,
"cumulative copy bytes",
),
] {
if exceeded {
return Err(ExecutionControlError::Limit(name));
}
}
state.usage = next;
Ok(SnapshotReservation {
lease: Rc::new(ReservationLease {
budget: self.clone(),
kind,
retained_bytes: estimate.retained_bytes,
}),
})
}
}
#[derive(Clone)]
pub struct SnapshotReservation {
lease: Rc<ReservationLease>,
}
impl SnapshotReservation {
pub fn retained_bytes(&self) -> u64 {
self.lease.retained_bytes
}
}
struct ReservationLease {
budget: SnapshotBudget,
kind: SnapshotResourceKind,
retained_bytes: u64,
}
impl Drop for ReservationLease {
fn drop(&mut self) {
let mut state = self.budget.0.borrow_mut();
state.usage.snapshots -= u64::from(self.kind == SnapshotResourceKind::Snapshot);
state.usage.branches -= u64::from(self.kind == SnapshotResourceKind::Branch);
state.usage.retained_bytes -= self.retained_bytes;
}
}
#[cfg(test)]
mod tests;