use crate::generation::GenerationCancellationToken;
use serde::{Deserialize, Serialize};
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
pub trait NativeTextStateBackend: crate::TextGenerationBackend {
type NativeTextState;
fn native_text_state_support(runtime: &crate::ModelRuntime<Self>) -> ControlSupport;
fn estimate_native_text_state(
runtime: &crate::ModelRuntime<Self>,
saved: Option<&Self::NativeTextState>,
) -> Result<Option<SnapshotEstimate>, Self::Error>;
fn estimate_native_text_growth(
_runtime: &crate::ModelRuntime<Self>,
_saved: &Self::NativeTextState,
_additional_input_tokens: u64,
) -> Result<Option<u64>, Self::Error> {
Ok(None)
}
fn capture_native_text_state(
runtime: &mut crate::ModelRuntime<Self>,
) -> Result<Self::NativeTextState, Self::Error>;
fn copy_native_text_state(
runtime: &mut crate::ModelRuntime<Self>,
saved: &Self::NativeTextState,
) -> Result<Self::NativeTextState, Self::Error>;
fn validate_native_text_state(
runtime: &crate::ModelRuntime<Self>,
saved: &Self::NativeTextState,
) -> Result<(), Self::Error>;
fn exchange_native_text_state(
runtime: &mut crate::ModelRuntime<Self>,
slot: &mut Self::NativeTextState,
) -> Result<(), Self::Error>;
}
pub const EXECUTION_CONTROL_SCHEMA_VERSION: u32 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum GenerationStatus {
Prepared,
Paused,
Running,
Completed,
Cancelled,
Failed,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "support", rename_all = "snake_case")]
pub enum ControlSupport {
Supported,
Unsupported {
reason: String,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SnapshotIsolation {
DeepCopy,
CopyOnWrite,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutionControlCapabilities {
pub schema_version: u32,
pub step: ControlSupport,
pub pause_resume: ControlSupport,
pub snapshot: ControlSupport,
pub restore: ControlSupport,
pub fork: ControlSupport,
pub force_next_token: ControlSupport,
pub sampling_overrides: ControlSupport,
pub isolation: Option<SnapshotIsolation>,
pub conditions: Vec<String>,
}
impl ExecutionControlCapabilities {
pub fn unsupported(reason: impl Into<String>) -> Self {
let support = ControlSupport::Unsupported {
reason: reason.into(),
};
Self {
schema_version: EXECUTION_CONTROL_SCHEMA_VERSION,
step: support.clone(),
pause_resume: support.clone(),
snapshot: support.clone(),
restore: support.clone(),
fork: support.clone(),
force_next_token: support.clone(),
sampling_overrides: support,
isolation: None,
conditions: Vec::new(),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct GenerationControlHandle {
pause: Arc<AtomicBool>,
cancellation: GenerationCancellationToken,
}
impl GenerationControlHandle {
pub fn new(cancellation: GenerationCancellationToken) -> Self {
Self {
pause: Arc::new(AtomicBool::new(false)),
cancellation,
}
}
pub fn request_pause(&self) {
self.pause.store(true, Ordering::Release);
}
pub fn pause_requested(&self) -> bool {
self.pause.load(Ordering::Acquire)
}
pub fn acknowledge_resume(&self) {
self.pause.store(false, Ordering::Release);
}
pub fn cancel(&self) {
self.cancellation.cancel();
}
pub fn cancellation(&self) -> &GenerationCancellationToken {
&self.cancellation
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct SnapshotEstimate {
pub retained_bytes: u64,
pub copy_bytes: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct SnapshotLimits {
pub max_snapshots: u64,
pub max_branches: u64,
pub retained_bytes: u64,
pub cumulative_copy_bytes: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum SnapshotResourceKind {
Snapshot,
Branch,
Restore,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct SnapshotUsage {
pub snapshots: u64,
pub branches: u64,
pub retained_bytes: u64,
pub cumulative_copy_bytes: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum ExecutionControlError {
#[error("invalid generation transition from {from:?} to {to:?}")]
Transition {
from: GenerationStatus,
to: GenerationStatus,
},
#[error("complete snapshot resource estimate is unavailable")]
UnknownEstimate,
#[error("execution-control accounting overflow")]
Overflow,
#[error("snapshot resource limit exceeded: {0}")]
Limit(&'static str),
}