use crate::{KvCacheHandle, RecurrentStateHandle, RecurrentStateSpec, TensorRef};
use async_trait::async_trait;
use ferrum_types::{ExecutorAdmissionLimits, FerrumError, ModelInfo, RequestId, Result, TokenId};
use serde::{Deserialize, Serialize};
use std::{
collections::{hash_map::DefaultHasher, HashMap, HashSet},
future::Future,
hash::{Hash, Hasher},
num::NonZeroU64,
ops::Range,
pin::Pin,
sync::Arc,
};
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KvSlotRequest {
pub cache_id: String,
pub target_len: usize,
pub admission_target_len: Option<usize>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KvSlotAllocation {
pub cache_id: String,
pub blocks_before: usize,
pub blocks_after: usize,
pub new_blocks: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KvSlotReservation {
pub block_size: usize,
pub total_blocks: usize,
pub free_blocks_before: usize,
pub free_blocks_after: usize,
pub allocations: Vec<KvSlotAllocation>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct KvSlotCapacitySnapshot {
pub block_size: usize,
pub total_blocks: usize,
pub free_blocks: usize,
}
#[derive(Clone)]
pub struct TokenSelectionMask {
pub fingerprint: u64,
pub valid_token_mask: Arc<[i8]>,
}
impl TokenSelectionMask {
pub fn new(valid_token_mask: Vec<i8>) -> Self {
let fingerprint = Self::fingerprint(&valid_token_mask);
Self {
fingerprint,
valid_token_mask: Arc::from(valid_token_mask),
}
}
fn fingerprint(valid_token_mask: &[i8]) -> u64 {
let mut hasher = DefaultHasher::new();
valid_token_mask.hash(&mut hasher);
hasher.finish()
}
pub fn set_tokens_validity(&mut self, token_ids: &[u32], valid: bool) -> bool {
let value = i8::from(valid);
let slots = Arc::make_mut(&mut self.valid_token_mask);
let mut changed = false;
for &token_id in token_ids {
if let Some(slot) = slots.get_mut(token_id as usize) {
if *slot != value {
*slot = value;
changed = true;
}
}
}
if changed {
self.fingerprint = Self::fingerprint(slots);
}
changed
}
pub fn len(&self) -> usize {
self.valid_token_mask.len()
}
pub fn is_empty(&self) -> bool {
self.valid_token_mask.is_empty()
}
}
impl std::fmt::Debug for TokenSelectionMask {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let valid_count = self.valid_token_mask.iter().filter(|&&v| v != 0).count();
f.debug_struct("TokenSelectionMask")
.field("fingerprint", &self.fingerprint)
.field("len", &self.valid_token_mask.len())
.field("valid_count", &valid_count)
.finish()
}
}
#[cfg(test)]
mod token_selection_mask_tests {
use super::TokenSelectionMask;
#[test]
fn response_completion_mask_is_copy_on_write_and_restores_fingerprint() {
let mut mask = TokenSelectionMask::new(vec![1, 1, 1]);
let original = mask.clone();
let original_fingerprint = mask.fingerprint;
assert!(mask.set_tokens_validity(&[1], false));
assert_eq!(mask.valid_token_mask.as_ref(), &[1, 0, 1]);
assert_eq!(original.valid_token_mask.as_ref(), &[1, 1, 1]);
assert_ne!(mask.fingerprint, original_fingerprint);
let masked_fingerprint = mask.fingerprint;
assert!(!mask.set_tokens_validity(&[1], false));
assert_eq!(mask.fingerprint, masked_fingerprint);
assert!(mask.set_tokens_validity(&[1], true));
assert_eq!(mask.fingerprint, original_fingerprint);
}
}
#[derive(Clone, Debug)]
pub enum LogitsReturnPolicy {
FullLogits,
GreedyArgmax {
token_mask: Option<TokenSelectionMask>,
repetition_penalty: Option<GreedyRepetitionPenalty>,
},
}
impl Default for LogitsReturnPolicy {
fn default() -> Self {
Self::FullLogits
}
}
impl LogitsReturnPolicy {
pub fn requires_full_logits(&self) -> bool {
matches!(self, Self::FullLogits)
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum ExecutorSamplingOutput {
FullLogits(Vec<f32>),
GreedyToken(TokenId),
}
impl ExecutorSamplingOutput {
pub fn full_logits(logits: Vec<f32>) -> Result<Self> {
if logits.is_empty() {
return Err(FerrumError::backend(
"plan-runtime sampling output requires non-empty logits",
));
}
Ok(Self::FullLogits(logits))
}
pub const fn greedy_token(token: TokenId) -> Self {
Self::GreedyToken(token)
}
pub fn validate_for_policy(
&self,
policy: &LogitsReturnPolicy,
vocabulary_size: usize,
) -> Result<()> {
match self {
Self::FullLogits(logits) if logits.len() != vocabulary_size => {
return Err(FerrumError::backend(format!(
"plan runtime returned {} logits for vocabulary {vocabulary_size}",
logits.len()
)));
}
Self::GreedyToken(_) if policy.requires_full_logits() => {
return Err(FerrumError::backend(
"plan runtime returned a greedy token for a full-logits request",
));
}
Self::GreedyToken(token)
if usize::try_from(token.get())
.ok()
.is_none_or(|token| token >= vocabulary_size) =>
{
return Err(FerrumError::backend(format!(
"plan runtime returned token {} outside vocabulary {vocabulary_size}",
token.get()
)));
}
_ => {}
}
Ok(())
}
pub fn into_full_logits(self) -> Result<Vec<f32>> {
match self {
Self::FullLogits(logits) => Ok(logits),
Self::GreedyToken(_) => Err(FerrumError::backend(
"plan-runtime prefill unexpectedly returned a selected token",
)),
}
}
}
#[cfg(test)]
mod executor_sampling_output_tests {
use super::{ExecutorSamplingOutput, LogitsReturnPolicy};
use ferrum_types::TokenId;
#[test]
fn full_logits_require_exact_vocabulary_width() {
let output = ExecutorSamplingOutput::full_logits(vec![0.0; 4]).unwrap();
assert!(output
.validate_for_policy(&LogitsReturnPolicy::FullLogits, 4)
.is_ok());
assert!(output
.validate_for_policy(&LogitsReturnPolicy::FullLogits, 5)
.is_err());
}
#[test]
fn greedy_token_requires_greedy_policy_and_in_vocabulary_token() {
let allowed = LogitsReturnPolicy::GreedyArgmax {
token_mask: None,
repetition_penalty: None,
};
let output = ExecutorSamplingOutput::greedy_token(TokenId::new(3));
assert!(output.validate_for_policy(&allowed, 4).is_ok());
assert!(output.validate_for_policy(&allowed, 3).is_err());
assert!(output
.validate_for_policy(&LogitsReturnPolicy::FullLogits, 4)
.is_err());
}
#[test]
fn full_logits_are_a_legal_greedy_batch_fallback() {
let policy = LogitsReturnPolicy::GreedyArgmax {
token_mask: None,
repetition_penalty: None,
};
let output = ExecutorSamplingOutput::full_logits(vec![0.0; 4]).unwrap();
assert!(output.validate_for_policy(&policy, 4).is_ok());
}
}
#[derive(Clone, Debug)]
pub struct GreedyRepetitionPenalty {
penalty: f32,
token_ids: Arc<[u32]>,
}
impl GreedyRepetitionPenalty {
pub fn new(penalty: f32, mut token_ids: Vec<u32>) -> Self {
let mut seen = HashSet::with_capacity(token_ids.len());
token_ids.retain(|token| seen.insert(*token));
Self {
penalty,
token_ids: Arc::from(token_ids),
}
}
pub const fn penalty(&self) -> f32 {
self.penalty
}
pub fn token_ids(&self) -> &[u32] {
&self.token_ids
}
pub fn is_empty(&self) -> bool {
self.token_ids.is_empty() || self.penalty == 1.0
}
}
#[cfg(test)]
mod greedy_repetition_penalty_tests {
use super::GreedyRepetitionPenalty;
#[test]
fn constructor_preserves_first_seen_order_and_removes_duplicates() {
let repetition = GreedyRepetitionPenalty::new(1.1, vec![7, 3, 7, 9, 3]);
assert_eq!(repetition.penalty(), 1.1);
assert_eq!(repetition.token_ids(), [7, 3, 9]);
}
}
#[derive(Debug, Clone)]
pub struct PrefillInput {
pub request_id: Option<RequestId>,
pub maximum_sequence_tokens: Option<usize>,
pub chunk: Option<PrefillChunk>,
pub input_ids: TensorRef,
pub attention_mask: Option<TensorRef>,
pub position_ids: Option<TensorRef>,
pub kv_cache: Option<Arc<dyn KvCacheHandle>>,
pub recurrent_state: Option<Arc<dyn RecurrentStateHandle>>,
pub metadata: HashMap<String, serde_json::Value>,
}
impl PrefillInput {
pub fn new(input_ids: TensorRef) -> Self {
Self {
request_id: None,
maximum_sequence_tokens: None,
chunk: None,
input_ids,
attention_mask: None,
position_ids: None,
kv_cache: None,
recurrent_state: None,
metadata: HashMap::new(),
}
}
pub fn with_request_context(
mut self,
request_id: RequestId,
maximum_sequence_tokens: usize,
) -> Self {
self.request_id = Some(request_id);
self.maximum_sequence_tokens = Some(maximum_sequence_tokens);
self
}
pub fn with_chunk(mut self, chunk: PrefillChunk) -> Self {
self.chunk = Some(chunk);
self
}
pub fn with_kv_cache(mut self, kv_cache: Arc<dyn KvCacheHandle>) -> Self {
self.kv_cache = Some(kv_cache);
self
}
pub fn with_recurrent_state(mut self, recurrent_state: Arc<dyn RecurrentStateHandle>) -> Self {
self.recurrent_state = Some(recurrent_state);
self
}
pub fn with_metadata(mut self, metadata: HashMap<String, serde_json::Value>) -> Self {
self.metadata = metadata;
self
}
pub fn with_attention_mask(mut self, mask: TensorRef) -> Self {
self.attention_mask = Some(mask);
self
}
pub fn with_position_ids(mut self, positions: TensorRef) -> Self {
self.position_ids = Some(positions);
self
}
pub fn batch_size(&self) -> usize {
self.input_ids.shape()[0]
}
pub fn sequence_length(&self) -> usize {
if self.input_ids.shape().len() >= 2 {
self.input_ids.shape()[1]
} else {
1
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
pub struct PrefillChunk {
tokens_processed: usize,
tokens_to_process: usize,
total_prompt_tokens: usize,
}
#[cfg(test)]
mod prefill_chunk_tests {
use super::PrefillChunk;
#[test]
fn validates_exact_progress_and_finality() {
let first = PrefillChunk::new(0, 3, 8).unwrap();
assert_eq!(first.range(), 0..3);
assert_eq!(first.end(), 3);
assert!(!first.is_final());
let final_chunk = PrefillChunk::new(3, 5, 8).unwrap();
assert_eq!(final_chunk.range(), 3..8);
assert!(final_chunk.is_final());
}
#[test]
fn rejects_empty_out_of_bounds_and_overflowing_progress() {
assert!(PrefillChunk::new(0, 0, 8).is_err());
assert!(PrefillChunk::new(0, 1, 0).is_err());
assert!(PrefillChunk::new(7, 2, 8).is_err());
assert!(PrefillChunk::new(usize::MAX, 1, usize::MAX).is_err());
}
}
impl PrefillChunk {
pub fn new(
tokens_processed: usize,
tokens_to_process: usize,
total_prompt_tokens: usize,
) -> Result<Self> {
let end = tokens_processed
.checked_add(tokens_to_process)
.ok_or_else(|| {
ferrum_types::FerrumError::request_validation("prefill chunk overflows")
})?;
if tokens_to_process == 0 || total_prompt_tokens == 0 || end > total_prompt_tokens {
return Err(ferrum_types::FerrumError::request_validation(
"prefill chunk must be non-empty and within the full prompt",
));
}
Ok(Self {
tokens_processed,
tokens_to_process,
total_prompt_tokens,
})
}
pub const fn tokens_processed(self) -> usize {
self.tokens_processed
}
pub const fn tokens_to_process(self) -> usize {
self.tokens_to_process
}
pub const fn total_prompt_tokens(self) -> usize {
self.total_prompt_tokens
}
pub fn range(self) -> Range<usize> {
self.tokens_processed..self.tokens_processed + self.tokens_to_process
}
pub const fn end(self) -> usize {
self.tokens_processed + self.tokens_to_process
}
pub const fn is_final(self) -> bool {
self.end() == self.total_prompt_tokens
}
}
#[derive(Debug, Clone)]
pub struct PrefillOutput {
pub logits: TensorRef,
pub kv_cache: Arc<dyn KvCacheHandle>,
pub recurrent_state: Option<Arc<dyn RecurrentStateHandle>>,
pub hidden_states: Option<Vec<TensorRef>>,
pub attention_weights: Option<Vec<TensorRef>>,
}
impl PrefillOutput {
pub fn new(logits: TensorRef, kv_cache: Arc<dyn KvCacheHandle>) -> Self {
Self {
logits,
kv_cache,
recurrent_state: None,
hidden_states: None,
attention_weights: None,
}
}
pub fn with_recurrent_state(mut self, recurrent_state: Arc<dyn RecurrentStateHandle>) -> Self {
self.recurrent_state = Some(recurrent_state);
self
}
pub fn last_token_logits(&self) -> Result<TensorRef> {
let shape = self.logits.shape();
if shape.len() != 3 {
return Err(ferrum_types::FerrumError::backend(
"Expected 3D logits tensor [batch, seq, vocab]",
));
}
let seq_len = shape[1];
if seq_len == 0 {
return Err(ferrum_types::FerrumError::backend("Empty sequence"));
}
self.logits
.view(&[0, seq_len - 1, 0], &[shape[0], seq_len, shape[2]])
}
}
#[derive(Debug, Clone)]
pub struct DecodeInput {
pub request_id: Option<RequestId>,
pub input_ids: TensorRef,
pub kv_cache: Arc<dyn KvCacheHandle>,
pub recurrent_state: Option<Arc<dyn RecurrentStateHandle>>,
pub position_ids: Option<TensorRef>,
pub metadata: HashMap<String, serde_json::Value>,
pub logits_policy: LogitsReturnPolicy,
}
impl DecodeInput {
pub fn new(input_ids: TensorRef, kv_cache: Arc<dyn KvCacheHandle>) -> Self {
Self {
request_id: None,
input_ids,
kv_cache,
recurrent_state: None,
position_ids: None,
metadata: HashMap::new(),
logits_policy: LogitsReturnPolicy::FullLogits,
}
}
pub fn with_request_id(mut self, request_id: RequestId) -> Self {
self.request_id = Some(request_id);
self
}
pub fn with_position_ids(mut self, positions: TensorRef) -> Self {
self.position_ids = Some(positions);
self
}
pub fn with_recurrent_state(mut self, recurrent_state: Arc<dyn RecurrentStateHandle>) -> Self {
self.recurrent_state = Some(recurrent_state);
self
}
pub fn with_metadata(mut self, metadata: HashMap<String, serde_json::Value>) -> Self {
self.metadata = metadata;
self
}
pub fn with_logits_policy(mut self, policy: LogitsReturnPolicy) -> Self {
self.logits_policy = policy;
self
}
pub fn batch_size(&self) -> usize {
self.input_ids.shape()[0]
}
}
#[derive(Clone)]
pub struct UnifiedBatchItem {
pub seq_id: String,
pub q_tokens: Vec<u32>,
pub kv_cache: Arc<dyn KvCacheHandle>,
pub recurrent_state: Option<Arc<dyn RecurrentStateHandle>>,
pub pos_offset: usize,
pub is_final_chunk: bool,
pub metadata: HashMap<String, serde_json::Value>,
pub logits_policy: LogitsReturnPolicy,
}
impl std::fmt::Debug for UnifiedBatchItem {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UnifiedBatchItem")
.field("seq_id", &self.seq_id)
.field("q_len", &self.q_tokens.len())
.field("has_recurrent_state", &self.recurrent_state.is_some())
.field("pos_offset", &self.pos_offset)
.field("is_final_chunk", &self.is_final_chunk)
.finish()
}
}
#[derive(Debug, Clone, Default)]
pub struct UnifiedBatch {
pub items: Vec<UnifiedBatchItem>,
}
impl UnifiedBatch {
pub fn new() -> Self {
Self::default()
}
pub fn total_q_tokens(&self) -> usize {
self.items.iter().map(|it| it.q_tokens.len()).sum()
}
pub fn num_sampled_items(&self) -> usize {
self.items.iter().filter(|it| it.is_final_chunk).count()
}
}
#[derive(Debug, Clone)]
pub struct PlanRuntimeDecodeInput {
pub request_id: RequestId,
pub input_token: TokenId,
pub kv_cache: Arc<dyn KvCacheHandle>,
pub logits_policy: LogitsReturnPolicy,
}
impl PlanRuntimeDecodeInput {
pub fn new(
request_id: RequestId,
input_token: TokenId,
kv_cache: Arc<dyn KvCacheHandle>,
) -> Self {
Self {
request_id,
input_token,
kv_cache,
logits_policy: LogitsReturnPolicy::FullLogits,
}
}
pub fn with_logits_policy(mut self, logits_policy: LogitsReturnPolicy) -> Self {
self.logits_policy = logits_policy;
self
}
}
#[derive(Debug, Clone)]
pub struct PlanRuntimePrefillInput {
pub request_id: RequestId,
pub input_tokens: Arc<[TokenId]>,
pub maximum_sequence_tokens: usize,
pub chunk: PrefillChunk,
}
impl PlanRuntimePrefillInput {
pub fn new(
request_id: RequestId,
input_tokens: impl Into<Arc<[TokenId]>>,
maximum_sequence_tokens: usize,
chunk: PrefillChunk,
) -> Result<Self> {
let input_tokens = input_tokens.into();
if input_tokens.is_empty() {
return Err(FerrumError::request_validation(
"plan-runtime prefill requires at least one input token",
));
}
if chunk.total_prompt_tokens() != input_tokens.len() {
return Err(FerrumError::request_validation(format!(
"plan-runtime prefill chunk declares {} prompt tokens for input length {}",
chunk.total_prompt_tokens(),
input_tokens.len()
)));
}
if maximum_sequence_tokens < input_tokens.len() {
return Err(FerrumError::request_validation(format!(
"plan-runtime sequence ceiling {maximum_sequence_tokens} does not cover prompt length {}",
input_tokens.len()
)));
}
Ok(Self {
request_id,
input_tokens,
maximum_sequence_tokens,
chunk,
})
}
}
#[derive(Debug, Clone)]
pub struct DecodeOutput {
pub logits: TensorRef,
pub kv_cache: Arc<dyn KvCacheHandle>,
pub recurrent_state: Option<Arc<dyn RecurrentStateHandle>>,
pub hidden_state: Option<TensorRef>,
pub attention_weights: Option<Vec<TensorRef>>,
}
impl DecodeOutput {
pub fn new(logits: TensorRef, kv_cache: Arc<dyn KvCacheHandle>) -> Self {
Self {
logits,
kv_cache,
recurrent_state: None,
hidden_state: None,
attention_weights: None,
}
}
pub fn with_recurrent_state(mut self, recurrent_state: Arc<dyn RecurrentStateHandle>) -> Self {
self.recurrent_state = Some(recurrent_state);
self
}
}
#[derive(Debug, Clone)]
pub struct PlanRuntimeDecodeOutput {
pub sampling_output: ExecutorSamplingOutput,
pub kv_cache: Arc<dyn KvCacheHandle>,
}
impl PlanRuntimeDecodeOutput {
pub fn new(sampling_output: ExecutorSamplingOutput, kv_cache: Arc<dyn KvCacheHandle>) -> Self {
Self {
sampling_output,
kv_cache,
}
}
}
#[derive(Debug)]
pub enum PlanRuntimePrefillProduct {
Intermediate,
FinalLogits(Vec<f32>),
}
#[derive(Debug)]
pub struct PlanRuntimePrefillAuthority {
request_id: RequestId,
committed_tokens: usize,
kv_cache: Arc<dyn KvCacheHandle>,
}
impl PlanRuntimePrefillAuthority {
pub fn request_id(&self) -> &RequestId {
&self.request_id
}
pub const fn committed_tokens(&self) -> usize {
self.committed_tokens
}
pub fn kv_cache(&self) -> &Arc<dyn KvCacheHandle> {
&self.kv_cache
}
pub fn into_cache(self) -> Arc<dyn KvCacheHandle> {
self.kv_cache
}
}
#[derive(Debug)]
pub struct PlanRuntimePrefillOutput {
authority: PlanRuntimePrefillAuthority,
product: PlanRuntimePrefillProduct,
}
impl PlanRuntimePrefillOutput {
pub fn intermediate(
request_id: RequestId,
committed_tokens: usize,
kv_cache: Arc<dyn KvCacheHandle>,
) -> Self {
Self {
authority: PlanRuntimePrefillAuthority {
request_id,
committed_tokens,
kv_cache,
},
product: PlanRuntimePrefillProduct::Intermediate,
}
}
pub fn final_logits(
request_id: RequestId,
committed_tokens: usize,
logits: Vec<f32>,
kv_cache: Arc<dyn KvCacheHandle>,
) -> Result<Self> {
if logits.is_empty() {
return Err(FerrumError::backend(
"plan-runtime final prefill returned empty logits",
));
}
Ok(Self {
authority: PlanRuntimePrefillAuthority {
request_id,
committed_tokens,
kv_cache,
},
product: PlanRuntimePrefillProduct::FinalLogits(logits),
})
}
pub fn request_id(&self) -> &RequestId {
self.authority.request_id()
}
pub const fn committed_tokens(&self) -> usize {
self.authority.committed_tokens()
}
pub fn product(&self) -> &PlanRuntimePrefillProduct {
&self.product
}
pub fn kv_cache(&self) -> &Arc<dyn KvCacheHandle> {
self.authority.kv_cache()
}
pub fn validate_for_completion(
&self,
expected_request_id: &RequestId,
completed_chunk: PrefillChunk,
vocabulary_size: usize,
) -> Result<()> {
if self.request_id() != expected_request_id {
return Err(FerrumError::backend(format!(
"plan runtime returned prefill output for request {}, expected {expected_request_id}",
self.request_id()
)));
}
if self.committed_tokens() != completed_chunk.end() {
return Err(FerrumError::backend(format!(
"plan runtime returned prefill extent {}, expected {}",
self.committed_tokens(),
completed_chunk.end()
)));
}
if self.kv_cache().num_tokens() != self.committed_tokens() {
return Err(FerrumError::backend(format!(
"plan runtime prefill cache `{}` reports {} tokens for committed extent {}",
self.kv_cache().cache_id(),
self.kv_cache().num_tokens(),
self.committed_tokens()
)));
}
match (&self.product, completed_chunk.is_final()) {
(PlanRuntimePrefillProduct::Intermediate, false) => Ok(()),
(PlanRuntimePrefillProduct::FinalLogits(logits), true)
if logits.len() == vocabulary_size =>
{
Ok(())
}
(PlanRuntimePrefillProduct::FinalLogits(logits), true) => {
Err(FerrumError::backend(format!(
"plan runtime returned {} final prefill logits for vocabulary {vocabulary_size}",
logits.len()
)))
}
(PlanRuntimePrefillProduct::Intermediate, true) => Err(FerrumError::backend(
"plan runtime returned an intermediate product for a final prefill chunk",
)),
(PlanRuntimePrefillProduct::FinalLogits(_), false) => Err(FerrumError::backend(
"plan runtime returned final logits for an intermediate prefill chunk",
)),
}
}
pub fn into_parts(self) -> (PlanRuntimePrefillAuthority, PlanRuntimePrefillProduct) {
(self.authority, self.product)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutorSequenceCompletion {
request_id: RequestId,
cache_id: String,
input_tokens: u64,
output_tokens: u64,
}
impl ExecutorSequenceCompletion {
pub fn new(
request_id: RequestId,
cache_id: String,
input_tokens: usize,
output_tokens: usize,
) -> Result<Self> {
if cache_id.is_empty() {
return Err(FerrumError::request_validation(
"executor sequence completion requires a cache identity",
));
}
let input_tokens = u64::try_from(input_tokens).map_err(|_| {
FerrumError::request_validation("executor completion input token count exceeds u64")
})?;
let output_tokens = u64::try_from(output_tokens).map_err(|_| {
FerrumError::request_validation("executor completion output token count exceeds u64")
})?;
Ok(Self {
request_id,
cache_id,
input_tokens,
output_tokens,
})
}
pub fn request_id(&self) -> &RequestId {
&self.request_id
}
pub fn cache_id(&self) -> &str {
&self.cache_id
}
pub const fn input_tokens(&self) -> u64 {
self.input_tokens
}
pub const fn output_tokens(&self) -> u64 {
self.output_tokens
}
}
pub use ferrum_types::ExecutionResourceAuthority;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutorExecutionCapacityPreemption {
request_id: RequestId,
cache_id: String,
}
impl ExecutorExecutionCapacityPreemption {
pub fn new(request_id: RequestId, cache_id: String) -> Result<Self> {
if cache_id.is_empty() {
return Err(FerrumError::request_validation(
"execution-capacity preemption requires a cache identity",
));
}
Ok(Self {
request_id,
cache_id,
})
}
pub fn request_id(&self) -> &RequestId {
&self.request_id
}
pub fn cache_id(&self) -> &str {
&self.cache_id
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ExecutorExecutionCapacityPreemptionAuthority {
RetainedPrefill,
ActiveSequence,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutorExecutionCapacityPreemptionReceipt {
request_id: RequestId,
cache_id: String,
authority: ExecutorExecutionCapacityPreemptionAuthority,
}
impl ExecutorExecutionCapacityPreemptionReceipt {
pub fn new(
request_id: RequestId,
cache_id: String,
authority: ExecutorExecutionCapacityPreemptionAuthority,
) -> Self {
Self {
request_id,
cache_id,
authority,
}
}
pub fn request_id(&self) -> &RequestId {
&self.request_id
}
pub fn cache_id(&self) -> &str {
&self.cache_id
}
pub const fn authority(&self) -> ExecutorExecutionCapacityPreemptionAuthority {
self.authority
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct PlanRuntimeResourceSnapshot {
device_capacity_bytes: u64,
usable_capacity_bytes: u64,
process_claimed_bytes: u64,
plan_claimed_bytes: u64,
static_bytes: u64,
dynamic_resident_bytes: u64,
dynamic_free_bytes: u64,
pending_growth_bytes: u64,
quarantined_bytes: u64,
}
impl PlanRuntimeResourceSnapshot {
#[allow(clippy::too_many_arguments)]
pub fn new(
device_capacity_bytes: u64,
usable_capacity_bytes: u64,
process_claimed_bytes: u64,
plan_claimed_bytes: u64,
static_bytes: u64,
dynamic_resident_bytes: u64,
dynamic_free_bytes: u64,
pending_growth_bytes: u64,
quarantined_bytes: u64,
) -> Result<Self> {
let snapshot = Self {
device_capacity_bytes,
usable_capacity_bytes,
process_claimed_bytes,
plan_claimed_bytes,
static_bytes,
dynamic_resident_bytes,
dynamic_free_bytes,
pending_growth_bytes,
quarantined_bytes,
};
snapshot.validate()?;
Ok(snapshot)
}
pub fn validate(&self) -> Result<()> {
if self.usable_capacity_bytes > self.device_capacity_bytes {
return Err(ferrum_types::FerrumError::internal(format!(
"plan runtime usable capacity {} exceeds device capacity {}",
self.usable_capacity_bytes, self.device_capacity_bytes
)));
}
if self.process_claimed_bytes > self.usable_capacity_bytes {
return Err(ferrum_types::FerrumError::internal(format!(
"plan runtime process claims {} exceed usable capacity {}",
self.process_claimed_bytes, self.usable_capacity_bytes
)));
}
if self.plan_claimed_bytes > self.process_claimed_bytes {
return Err(ferrum_types::FerrumError::internal(format!(
"plan runtime plan claims {} exceed process claims {}",
self.plan_claimed_bytes, self.process_claimed_bytes
)));
}
if self.dynamic_free_bytes > self.dynamic_resident_bytes {
return Err(ferrum_types::FerrumError::internal(format!(
"plan runtime dynamic free bytes {} exceed resident bytes {}",
self.dynamic_free_bytes, self.dynamic_resident_bytes
)));
}
let minimum_plan_claim = self
.static_bytes
.checked_add(self.dynamic_resident_bytes)
.and_then(|bytes| bytes.checked_add(self.quarantined_bytes))
.ok_or_else(|| {
ferrum_types::FerrumError::internal(
"plan runtime static, resident, and quarantined bytes overflow u64",
)
})?;
if minimum_plan_claim > self.plan_claimed_bytes {
return Err(ferrum_types::FerrumError::internal(format!(
"plan runtime accounted plan bytes {minimum_plan_claim} exceed plan claims {}",
self.plan_claimed_bytes
)));
}
Ok(())
}
pub const fn device_capacity_bytes(&self) -> u64 {
self.device_capacity_bytes
}
pub const fn usable_capacity_bytes(&self) -> u64 {
self.usable_capacity_bytes
}
pub const fn process_claimed_bytes(&self) -> u64 {
self.process_claimed_bytes
}
pub const fn plan_claimed_bytes(&self) -> u64 {
self.plan_claimed_bytes
}
pub const fn static_bytes(&self) -> u64 {
self.static_bytes
}
pub const fn dynamic_resident_bytes(&self) -> u64 {
self.dynamic_resident_bytes
}
pub const fn dynamic_free_bytes(&self) -> u64 {
self.dynamic_free_bytes
}
pub const fn dynamic_used_bytes(&self) -> u64 {
self.dynamic_resident_bytes - self.dynamic_free_bytes
}
pub const fn pending_growth_bytes(&self) -> u64 {
self.pending_growth_bytes
}
pub const fn quarantined_bytes(&self) -> u64 {
self.quarantined_bytes
}
pub fn available_bytes(&self) -> Result<u64> {
self.usable_capacity_bytes
.checked_sub(self.process_claimed_bytes)
.and_then(|bytes| bytes.checked_add(self.dynamic_free_bytes))
.ok_or_else(|| {
ferrum_types::FerrumError::internal(
"plan runtime available capacity calculation overflowed",
)
})
}
pub fn used_bytes(&self) -> Result<u64> {
self.available_bytes().and_then(|available| {
self.usable_capacity_bytes
.checked_sub(available)
.ok_or_else(|| {
ferrum_types::FerrumError::internal(
"plan runtime available bytes exceed usable capacity",
)
})
})
}
}
#[cfg(test)]
mod plan_runtime_resource_snapshot_tests {
use super::PlanRuntimeResourceSnapshot;
#[test]
fn separates_static_and_dynamic_usage() {
let snapshot =
PlanRuntimeResourceSnapshot::new(1_000, 900, 710, 710, 400, 300, 200, 20, 10).unwrap();
assert_eq!(snapshot.available_bytes().unwrap(), 390);
assert_eq!(snapshot.used_bytes().unwrap(), 510);
assert_eq!(snapshot.dynamic_resident_bytes(), 300);
assert_eq!(snapshot.dynamic_used_bytes(), 100);
assert_eq!(snapshot.dynamic_free_bytes(), 200);
assert_eq!(snapshot.pending_growth_bytes(), 20);
assert_eq!(snapshot.quarantined_bytes(), 10);
}
#[test]
fn rejects_incoherent_capacity_evidence() {
assert!(PlanRuntimeResourceSnapshot::new(1_000, 1_001, 0, 0, 0, 0, 0, 0, 0).is_err());
assert!(PlanRuntimeResourceSnapshot::new(1_000, 900, 901, 0, 0, 0, 0, 0, 0).is_err());
assert!(PlanRuntimeResourceSnapshot::new(1_000, 900, 500, 501, 0, 0, 0, 0, 0).is_err());
assert!(PlanRuntimeResourceSnapshot::new(1_000, 900, 100, 100, 0, 100, 101, 0, 0).is_err());
assert!(PlanRuntimeResourceSnapshot::new(1_000, 900, 500, 500, 400, 100, 0, 0, 1).is_err());
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ExecutorRequestOrigin {
Product,
Startup,
Diagnostic,
}
impl ExecutorRequestOrigin {
pub const fn namespace(self) -> &'static str {
match self {
Self::Product => "product",
Self::Startup => "startup",
Self::Diagnostic => "diagnostic",
}
}
pub fn from_namespaced_request_identity(identity: &str) -> Option<Self> {
let suffix = identity.strip_prefix("request.")?;
let (namespace, request_id) = suffix.split_once('.')?;
if request_id.is_empty() {
return None;
}
match namespace {
"product" => Some(Self::Product),
"startup" => Some(Self::Startup),
"diagnostic" => Some(Self::Diagnostic),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct ExecutorPrefillAdmission<'a> {
pub request_id: &'a RequestId,
pub input_tokens: &'a [TokenId],
pub maximum_sequence_tokens: usize,
pub product_prompt_tokens: usize,
pub replayed_output_tokens: usize,
pub request_origin: ExecutorRequestOrigin,
}
impl<'a> ExecutorPrefillAdmission<'a> {
pub const fn for_startup(
request_id: &'a RequestId,
input_tokens: &'a [TokenId],
maximum_sequence_tokens: usize,
) -> Self {
Self {
request_id,
input_tokens,
maximum_sequence_tokens,
product_prompt_tokens: input_tokens.len(),
replayed_output_tokens: 0,
request_origin: ExecutorRequestOrigin::Startup,
}
}
pub const fn for_diagnostic(
request_id: &'a RequestId,
input_tokens: &'a [TokenId],
maximum_sequence_tokens: usize,
) -> Self {
Self {
request_id,
input_tokens,
maximum_sequence_tokens,
product_prompt_tokens: input_tokens.len(),
replayed_output_tokens: 0,
request_origin: ExecutorRequestOrigin::Diagnostic,
}
}
pub fn for_product_request(
request_id: &'a RequestId,
input_tokens: &'a [TokenId],
maximum_sequence_tokens: usize,
product_prompt_tokens: usize,
replayed_output_tokens: usize,
) -> Result<Self> {
let admission = Self {
request_id,
input_tokens,
maximum_sequence_tokens,
product_prompt_tokens,
replayed_output_tokens,
request_origin: ExecutorRequestOrigin::Product,
};
admission.validate()?;
Ok(admission)
}
pub fn validate(&self) -> Result<()> {
if self.input_tokens.is_empty() {
return Err(FerrumError::request_validation(
"executor prefill admission requires at least one execution-context token",
));
}
if self.product_prompt_tokens == 0 {
return Err(FerrumError::request_validation(
"executor prefill admission requires at least one product prompt token",
));
}
let execution_context_tokens = self
.product_prompt_tokens
.checked_add(self.replayed_output_tokens)
.ok_or_else(|| {
FerrumError::request_validation(
"executor prefill product token accounting exceeds usize",
)
})?;
if execution_context_tokens != self.input_tokens.len() {
return Err(FerrumError::request_validation(format!(
"executor prefill execution context has {} tokens but product accounting declares {} prompt + {} replayed output",
self.input_tokens.len(),
self.product_prompt_tokens,
self.replayed_output_tokens
)));
}
if self.maximum_sequence_tokens < execution_context_tokens {
return Err(FerrumError::request_validation(format!(
"executor prefill sequence ceiling {} does not cover execution context {execution_context_tokens}",
self.maximum_sequence_tokens
)));
}
Ok(())
}
}
#[cfg(test)]
mod executor_prefill_admission_tests {
use super::{ExecutorPrefillAdmission, ExecutorRequestOrigin};
use ferrum_types::{RequestId, TokenId};
#[test]
fn product_accounting_distinguishes_replayed_output_from_prompt() {
let request_id = RequestId::new();
let tokens = [1, 2, 3, 4, 5]
.into_iter()
.map(TokenId::new)
.collect::<Vec<_>>();
let admission =
ExecutorPrefillAdmission::for_product_request(&request_id, &tokens, 8, 3, 2)
.expect("recompute accounting must be accepted");
assert_eq!(admission.product_prompt_tokens, 3);
assert_eq!(admission.replayed_output_tokens, 2);
assert_eq!(admission.request_origin, ExecutorRequestOrigin::Product);
assert_eq!(
ExecutorPrefillAdmission::for_startup(&request_id, &tokens, 8).request_origin,
ExecutorRequestOrigin::Startup
);
assert_eq!(
ExecutorPrefillAdmission::for_diagnostic(&request_id, &tokens, 8).request_origin,
ExecutorRequestOrigin::Diagnostic
);
assert_eq!(ExecutorRequestOrigin::Product.namespace(), "product");
assert_eq!(ExecutorRequestOrigin::Startup.namespace(), "startup");
assert_eq!(ExecutorRequestOrigin::Diagnostic.namespace(), "diagnostic");
assert_eq!(
ExecutorRequestOrigin::from_namespaced_request_identity("request.product.123"),
Some(ExecutorRequestOrigin::Product)
);
assert_eq!(
ExecutorRequestOrigin::from_namespaced_request_identity("request.startup.123"),
Some(ExecutorRequestOrigin::Startup)
);
assert_eq!(
ExecutorRequestOrigin::from_namespaced_request_identity("request.diagnostic.123"),
Some(ExecutorRequestOrigin::Diagnostic)
);
assert_eq!(
ExecutorRequestOrigin::from_namespaced_request_identity("request.product."),
None
);
assert_eq!(
ExecutorRequestOrigin::from_namespaced_request_identity("request/external"),
None
);
}
#[test]
fn product_accounting_rejects_context_drift_and_short_ceiling() {
let request_id = RequestId::new();
let tokens = [1, 2, 3].into_iter().map(TokenId::new).collect::<Vec<_>>();
assert!(
ExecutorPrefillAdmission::for_product_request(&request_id, &tokens, 3, 2, 0).is_err()
);
assert!(
ExecutorPrefillAdmission::for_product_request(&request_id, &tokens, 2, 2, 1).is_err()
);
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct ExecutorPrefillAdmissionReceipt {
pub request_id: RequestId,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
pub struct ExecutorAdmissionEpochs {
pub coordinator_id: NonZeroU64,
pub release_epoch: u64,
pub capacity_epoch: u64,
}
impl ExecutorAdmissionEpochs {
pub const fn new(coordinator_id: NonZeroU64, release_epoch: u64, capacity_epoch: u64) -> Self {
Self {
coordinator_id,
release_epoch,
capacity_epoch,
}
}
pub fn from_capacity(epochs: crate::vnext::CapacityEpochs) -> Self {
Self::new(
NonZeroU64::new(epochs.coordinator_id().get())
.expect("core-issued admission coordinator ids are non-zero"),
epochs.release_epoch(),
epochs.capacity_epoch(),
)
}
}
type ExecutorCapacityWaitFuture =
Pin<Box<dyn Future<Output = Result<ExecutorAdmissionEpochs>> + Send + 'static>>;
#[must_use = "capacity wait registrations must be awaited or explicitly dropped"]
pub struct ExecutorCapacityWaitRegistration {
future: ExecutorCapacityWaitFuture,
}
impl ExecutorCapacityWaitRegistration {
pub fn new<F>(future: F) -> Self
where
F: Future<Output = Result<ExecutorAdmissionEpochs>> + Send + 'static,
{
Self {
future: Box::pin(future),
}
}
pub async fn wait_for_change(self) -> Result<ExecutorAdmissionEpochs> {
self.future.await
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ExecutorExecutionCapacityStage {
SequenceExtension,
StepAdmission,
SubmissionWave,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct ExecutorExecutionMaintenanceMutation {
pool_id: crate::vnext::DynamicBackingPoolId,
domain_id: crate::vnext::CapacityDomainId,
chunk: crate::vnext::BackingChunkIdentity,
chunk_bytes: u64,
published_capacity_bytes: u64,
capacity_epoch: u64,
}
impl ExecutorExecutionMaintenanceMutation {
pub fn pool_id(&self) -> &crate::vnext::DynamicBackingPoolId {
&self.pool_id
}
pub const fn domain_id(&self) -> crate::vnext::CapacityDomainId {
self.domain_id
}
pub fn chunk(&self) -> &crate::vnext::BackingChunkIdentity {
&self.chunk
}
pub const fn chunk_bytes(&self) -> u64 {
self.chunk_bytes
}
pub const fn published_capacity_bytes(&self) -> u64 {
self.published_capacity_bytes
}
pub const fn capacity_epoch(&self) -> u64 {
self.capacity_epoch
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct ExecutorExecutionMaintenanceProgress {
attempts: u32,
coordinator_id: NonZeroU64,
mutations: Vec<ExecutorExecutionMaintenanceMutation>,
latest_capacity_epoch: u64,
}
impl ExecutorExecutionMaintenanceProgress {
pub fn from_growth_receipts(
attempts: u32,
observed: ExecutorAdmissionEpochs,
receipts: &[crate::vnext::DynamicPoolGrowthBatchReceipt],
pools: &[crate::vnext::DynamicPoolStatus],
) -> Result<Self> {
if attempts == 0 || receipts.is_empty() || receipts.len() > attempts as usize {
return Err(FerrumError::internal(
"execution maintenance retry requires bounded, non-empty growth receipts",
));
}
let mut mutations = Vec::new();
let mut previous_capacity_epoch = None;
for receipt in receipts {
if receipt.coordinator_id().get() != observed.coordinator_id.get() {
return Err(FerrumError::internal(
"execution maintenance receipt belongs to another capacity coordinator",
));
}
if receipt.growths().is_empty()
|| previous_capacity_epoch
.is_some_and(|previous| receipt.capacity_epoch() <= previous)
{
return Err(FerrumError::internal(
"execution maintenance receipts contain no new ordered capacity mutation",
));
}
previous_capacity_epoch = Some(receipt.capacity_epoch());
for growth in receipt.growths() {
let pool = pools
.iter()
.find(|pool| pool.pool_id() == growth.pool_id())
.ok_or_else(|| {
FerrumError::internal(
"execution maintenance receipt references an unknown dynamic pool",
)
})?;
if growth.chunk().pool_id() != growth.pool_id()
|| growth.chunk_bytes() == 0
|| growth.published_capacity_bytes() == 0
|| growth.capacity_epoch() != receipt.capacity_epoch()
{
return Err(FerrumError::internal(
"execution maintenance receipt contains an invalid pool mutation",
));
}
if mutations
.iter()
.any(|mutation: &ExecutorExecutionMaintenanceMutation| {
mutation.pool_id() == growth.pool_id() && mutation.chunk() == growth.chunk()
})
{
return Err(FerrumError::internal(
"execution maintenance receipts repeat one physical pool mutation",
));
}
mutations.push(ExecutorExecutionMaintenanceMutation {
pool_id: growth.pool_id().clone(),
domain_id: pool.domain_id(),
chunk: growth.chunk().clone(),
chunk_bytes: growth.chunk_bytes(),
published_capacity_bytes: growth.published_capacity_bytes(),
capacity_epoch: growth.capacity_epoch(),
});
}
}
let latest_capacity_epoch = previous_capacity_epoch.expect("receipts are non-empty");
if latest_capacity_epoch > observed.capacity_epoch {
return Err(FerrumError::internal(
"execution maintenance receipt is newer than the exported capacity observation",
));
}
mutations.sort_by(|left, right| {
(
left.capacity_epoch,
left.pool_id.as_str(),
left.chunk.ordinal(),
left.chunk.generation(),
)
.cmp(&(
right.capacity_epoch,
right.pool_id.as_str(),
right.chunk.ordinal(),
right.chunk.generation(),
))
});
Ok(Self {
attempts,
coordinator_id: observed.coordinator_id,
mutations,
latest_capacity_epoch,
})
}
pub const fn attempts(&self) -> u32 {
self.attempts
}
pub const fn coordinator_id(&self) -> NonZeroU64 {
self.coordinator_id
}
pub fn mutations(&self) -> &[ExecutorExecutionMaintenanceMutation] {
&self.mutations
}
pub const fn latest_capacity_epoch(&self) -> u64 {
self.latest_capacity_epoch
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct ExecutorExecutionMaintenanceRetry {
affected_request_ids: Vec<RequestId>,
progress: ExecutorExecutionMaintenanceProgress,
}
impl ExecutorExecutionMaintenanceRetry {
fn new(
affected_request_ids: Vec<RequestId>,
progress: ExecutorExecutionMaintenanceProgress,
) -> Result<Self> {
let unique = affected_request_ids.iter().collect::<HashSet<_>>();
if affected_request_ids.is_empty() || unique.len() != affected_request_ids.len() {
return Err(FerrumError::internal(
"execution maintenance retry requires unique affected requests",
));
}
if progress.mutations().is_empty() {
return Err(FerrumError::internal(
"execution maintenance retry requires physical mutations",
));
}
Ok(Self {
affected_request_ids,
progress,
})
}
pub fn affected_request_ids(&self) -> &[RequestId] {
&self.affected_request_ids
}
pub const fn progress(&self) -> &ExecutorExecutionMaintenanceProgress {
&self.progress
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ExecutorExecutionCapacityEvidenceOwner {
Logical,
Backing,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
enum ExecutorExecutionCapacityEvidenceKind {
Logical {
shortfalls: Vec<crate::vnext::CapacityShortfall>,
#[serde(skip_serializing_if = "Option::is_none")]
pressure: Option<crate::vnext::DynamicBackingPressure>,
},
BackingDeferred {
blockers: Vec<crate::vnext::DynamicBackingBlocker>,
#[serde(skip_serializing_if = "Option::is_none")]
pressure: Option<crate::vnext::DynamicBackingPressure>,
},
BackingPressure {
pressure: crate::vnext::DynamicBackingPressure,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct ExecutorExecutionCapacityEvidence {
owner: ExecutorExecutionCapacityEvidenceOwner,
#[serde(flatten)]
kind: ExecutorExecutionCapacityEvidenceKind,
#[serde(skip_serializing_if = "Option::is_none")]
maintenance_boundary: Option<crate::vnext::DynamicPoolMaintenanceBoundaryReceipt>,
}
impl ExecutorExecutionCapacityEvidence {
fn logical(shortfalls: Vec<crate::vnext::CapacityShortfall>) -> Result<Self> {
Self::logical_with_pressure(shortfalls, None)
}
fn logical_with_pressure(
shortfalls: Vec<crate::vnext::CapacityShortfall>,
pressure: Option<crate::vnext::DynamicBackingPressure>,
) -> Result<Self> {
if shortfalls.is_empty() {
return Err(FerrumError::internal(
"logical execution deferral requires at least one shortfall",
));
}
Ok(Self {
owner: ExecutorExecutionCapacityEvidenceOwner::Logical,
kind: ExecutorExecutionCapacityEvidenceKind::Logical {
shortfalls,
pressure,
},
maintenance_boundary: None,
})
}
fn backing_deferred(blockers: Vec<crate::vnext::DynamicBackingBlocker>) -> Result<Self> {
Self::backing_deferred_with_pressure(blockers, None)
}
fn backing_deferred_with_pressure(
blockers: Vec<crate::vnext::DynamicBackingBlocker>,
pressure: Option<crate::vnext::DynamicBackingPressure>,
) -> Result<Self> {
if blockers.is_empty() {
return Err(FerrumError::internal(
"physical execution deferral requires at least one backing blocker",
));
}
Ok(Self {
owner: ExecutorExecutionCapacityEvidenceOwner::Backing,
kind: ExecutorExecutionCapacityEvidenceKind::BackingDeferred { blockers, pressure },
maintenance_boundary: None,
})
}
fn direct_backing_pressure(pressure: crate::vnext::DynamicBackingPressure) -> Self {
Self {
owner: ExecutorExecutionCapacityEvidenceOwner::Backing,
kind: ExecutorExecutionCapacityEvidenceKind::BackingPressure { pressure },
maintenance_boundary: None,
}
}
pub const fn owner(&self) -> ExecutorExecutionCapacityEvidenceOwner {
self.owner
}
pub fn shortfalls(&self) -> &[crate::vnext::CapacityShortfall] {
match &self.kind {
ExecutorExecutionCapacityEvidenceKind::Logical { shortfalls, .. } => shortfalls,
ExecutorExecutionCapacityEvidenceKind::BackingDeferred { .. }
| ExecutorExecutionCapacityEvidenceKind::BackingPressure { .. } => &[],
}
}
pub fn backing_blockers(&self) -> &[crate::vnext::DynamicBackingBlocker] {
match &self.kind {
ExecutorExecutionCapacityEvidenceKind::BackingDeferred { blockers, .. } => blockers,
ExecutorExecutionCapacityEvidenceKind::Logical { .. }
| ExecutorExecutionCapacityEvidenceKind::BackingPressure { .. } => &[],
}
}
pub const fn backing_pressure(&self) -> Option<&crate::vnext::DynamicBackingPressure> {
match &self.kind {
ExecutorExecutionCapacityEvidenceKind::Logical { pressure, .. }
| ExecutorExecutionCapacityEvidenceKind::BackingDeferred { pressure, .. } => {
pressure.as_ref()
}
ExecutorExecutionCapacityEvidenceKind::BackingPressure { pressure } => Some(pressure),
}
}
pub const fn maintenance_boundary(
&self,
) -> Option<&crate::vnext::DynamicPoolMaintenanceBoundaryReceipt> {
self.maintenance_boundary.as_ref()
}
fn with_maintenance_boundary(
mut self,
boundary: Option<crate::vnext::DynamicPoolMaintenanceBoundaryReceipt>,
) -> Result<Self> {
match (self.backing_pressure(), boundary.as_ref()) {
(
Some(crate::vnext::DynamicBackingPressure::DeviceCapacity(pressure)),
Some(boundary),
) if pressure == boundary.pressure() && !boundary.reclaim_sufficient() => {}
(Some(crate::vnext::DynamicBackingPressure::PoolResident(_)), None) | (None, None) => {}
(Some(crate::vnext::DynamicBackingPressure::DeviceCapacity(_)), None) => {
return Err(FerrumError::internal(
"device-capacity execution maintenance lost its boundary receipt",
));
}
_ => {
return Err(FerrumError::internal(
"execution maintenance boundary differs from its blocked pressure",
));
}
}
self.maintenance_boundary = boundary;
Ok(self)
}
fn has_relevant_mutation(&self, mutation: &ExecutorExecutionMaintenanceMutation) -> bool {
let logical_matches = |shortfalls: &[crate::vnext::CapacityShortfall]| {
shortfalls.iter().any(|shortfall| {
shortfall.kind() == crate::vnext::CapacityShortfallKind::BackingGrowthRequired
&& shortfall.domain() == Some(mutation.domain_id())
})
};
let backing_matches = |blockers: &[crate::vnext::DynamicBackingBlocker]| {
blockers.iter().any(|blocker| {
blocker.pool_id() == mutation.pool_id()
&& blocker.domain_id() == mutation.domain_id()
})
};
match &self.kind {
ExecutorExecutionCapacityEvidenceKind::Logical { shortfalls, .. } => {
logical_matches(shortfalls)
}
ExecutorExecutionCapacityEvidenceKind::BackingDeferred { blockers, .. } => {
backing_matches(blockers)
}
ExecutorExecutionCapacityEvidenceKind::BackingPressure { .. } => false,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct ExecutorExecutionCapacityDeferral {
observed: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
stage: ExecutorExecutionCapacityStage,
evidence: ExecutorExecutionCapacityEvidence,
#[serde(skip_serializing_if = "Option::is_none")]
maintenance_retry: Option<ExecutorExecutionMaintenanceRetry>,
}
impl ExecutorExecutionCapacityDeferral {
fn with_evidence(
observed: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
stage: ExecutorExecutionCapacityStage,
evidence: ExecutorExecutionCapacityEvidence,
) -> Result<Self> {
if wait_condition.coordinator_id().get() != observed.coordinator_id.get() {
return Err(ferrum_types::FerrumError::request_validation(
"executor execution deferral belongs to a different capacity coordinator",
));
}
Ok(Self {
observed,
wait_condition,
stage,
evidence,
maintenance_retry: None,
})
}
pub fn from_backing_pressure(
observed: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
pressure: crate::vnext::DynamicBackingPressure,
stage: ExecutorExecutionCapacityStage,
) -> Result<Self> {
Self::with_evidence(
observed,
wait_condition,
stage,
ExecutorExecutionCapacityEvidence::direct_backing_pressure(pressure),
)
}
pub fn from_admission(
deferred: &crate::vnext::AdmissionDeferred,
stage: ExecutorExecutionCapacityStage,
) -> Result<Self> {
if deferred.action() != crate::vnext::DeferredAction::WaitForRelease {
return Err(ferrum_types::FerrumError::internal(
"execution capacity deferral must be reduced to WaitForRelease before export",
));
}
let evidence = ExecutorExecutionCapacityEvidence::logical(deferred.blockers().to_vec())?;
Self::with_evidence(
ExecutorAdmissionEpochs::from_capacity(deferred.epochs()),
deferred.wait_condition().clone(),
stage,
evidence,
)
}
pub fn from_pending_maintenance(
deferred: &crate::vnext::AdmissionDeferred,
stage: ExecutorExecutionCapacityStage,
) -> Result<Self> {
if deferred.action() != crate::vnext::DeferredAction::AwaitBackingGrowth {
return Err(ferrum_types::FerrumError::internal(
"pending execution maintenance must await backing growth",
));
}
let evidence = ExecutorExecutionCapacityEvidence::logical(deferred.blockers().to_vec())?;
Self::with_evidence(
ExecutorAdmissionEpochs::from_capacity(deferred.epochs()),
deferred.wait_condition().clone(),
stage,
evidence,
)
}
pub fn from_backing(
deferred: &crate::vnext::DynamicBackingDeferred,
stage: ExecutorExecutionCapacityStage,
) -> Result<Self> {
let evidence =
ExecutorExecutionCapacityEvidence::backing_deferred(deferred.blockers().to_vec())?;
Self::with_evidence(
ExecutorAdmissionEpochs::from_capacity(deferred.epochs()),
deferred.wait_condition().clone(),
stage,
evidence,
)
}
pub fn with_relevant_maintenance_retry(
mut self,
attempts: u32,
receipts: &[crate::vnext::DynamicPoolGrowthBatchReceipt],
pools: &[crate::vnext::DynamicPoolStatus],
affected_request_ids: Vec<RequestId>,
) -> Result<Self> {
if receipts.is_empty() {
return Ok(self);
}
let progress = ExecutorExecutionMaintenanceProgress::from_growth_receipts(
attempts,
self.observed,
receipts,
pools,
)?;
if progress.coordinator_id() != self.observed.coordinator_id
|| progress.latest_capacity_epoch() > self.observed.capacity_epoch
|| progress.mutations().is_empty()
{
return Err(FerrumError::internal(
"execution maintenance progress does not match the exported deferral",
));
}
let relevant_mutation = progress
.mutations()
.iter()
.any(|mutation| self.evidence.has_relevant_mutation(mutation));
if !relevant_mutation {
return Ok(self);
}
let retry = ExecutorExecutionMaintenanceRetry::new(affected_request_ids, progress)?;
if self.stage == ExecutorExecutionCapacityStage::SequenceExtension
&& retry.affected_request_ids().len() != 1
{
return Err(FerrumError::internal(
"sequence-extension maintenance retry must affect exactly one request",
));
}
self.maintenance_retry = Some(retry);
Ok(self)
}
pub fn from_admission_maintenance(
source: &crate::vnext::AdmissionDeferred,
observed: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
pressure: crate::vnext::DynamicBackingPressure,
maintenance_boundary: Option<crate::vnext::DynamicPoolMaintenanceBoundaryReceipt>,
stage: ExecutorExecutionCapacityStage,
) -> Result<Self> {
if source.action() != crate::vnext::DeferredAction::AwaitBackingGrowth {
return Err(ferrum_types::FerrumError::internal(
"execution maintenance source must await backing growth",
));
}
let evidence = ExecutorExecutionCapacityEvidence::logical_with_pressure(
source.blockers().to_vec(),
Some(pressure),
)?
.with_maintenance_boundary(maintenance_boundary)?;
if evidence.maintenance_boundary().is_some_and(|boundary| {
boundary.coordinator_id().get() != observed.coordinator_id.get()
}) {
return Err(FerrumError::internal(
"execution maintenance boundary belongs to another coordinator",
));
}
Self::with_evidence(observed, wait_condition, stage, evidence)
}
pub fn from_backing_maintenance(
source: &crate::vnext::DynamicBackingDeferred,
observed: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
pressure: crate::vnext::DynamicBackingPressure,
maintenance_boundary: Option<crate::vnext::DynamicPoolMaintenanceBoundaryReceipt>,
stage: ExecutorExecutionCapacityStage,
) -> Result<Self> {
let evidence = ExecutorExecutionCapacityEvidence::backing_deferred_with_pressure(
source.blockers().to_vec(),
Some(pressure),
)?
.with_maintenance_boundary(maintenance_boundary)?;
if evidence.maintenance_boundary().is_some_and(|boundary| {
boundary.coordinator_id().get() != observed.coordinator_id.get()
}) {
return Err(FerrumError::internal(
"execution maintenance boundary belongs to another coordinator",
));
}
Self::with_evidence(observed, wait_condition, stage, evidence)
}
pub const fn observed(&self) -> ExecutorAdmissionEpochs {
self.observed
}
pub fn wait_condition(&self) -> &crate::vnext::CapacityWaitCondition {
&self.wait_condition
}
pub const fn stage(&self) -> ExecutorExecutionCapacityStage {
self.stage
}
pub const fn evidence(&self) -> &ExecutorExecutionCapacityEvidence {
&self.evidence
}
pub fn shortfalls(&self) -> &[crate::vnext::CapacityShortfall] {
self.evidence.shortfalls()
}
pub fn backing_blockers(&self) -> &[crate::vnext::DynamicBackingBlocker] {
self.evidence.backing_blockers()
}
pub const fn backing_pressure(&self) -> Option<&crate::vnext::DynamicBackingPressure> {
self.evidence.backing_pressure()
}
pub const fn maintenance_boundary(
&self,
) -> Option<&crate::vnext::DynamicPoolMaintenanceBoundaryReceipt> {
self.evidence.maintenance_boundary()
}
pub fn maintenance_retry(&self) -> Option<&ExecutorExecutionMaintenanceRetry> {
self.maintenance_retry.as_ref()
}
pub fn validated_maintenance_retry_scope(
&self,
current_request_ids: &[RequestId],
) -> Result<Option<&ExecutorExecutionMaintenanceRetry>> {
let Some(retry) = self.maintenance_retry.as_ref() else {
return Ok(None);
};
let current = current_request_ids.iter().collect::<HashSet<_>>();
if current_request_ids.is_empty() || current.len() != current_request_ids.len() {
return Err(FerrumError::internal(
"execution maintenance retry received an invalid current request cohort",
));
}
let affected = retry.affected_request_ids().iter().collect::<HashSet<_>>();
if !affected.is_subset(¤t) {
return Err(FerrumError::internal(
"execution maintenance retry affects a request outside the current cohort",
));
}
match self.stage {
ExecutorExecutionCapacityStage::SequenceExtension => {
if affected.len() != 1 {
return Err(FerrumError::internal(
"sequence-extension maintenance retry must affect one current request",
));
}
}
ExecutorExecutionCapacityStage::StepAdmission
| ExecutorExecutionCapacityStage::SubmissionWave => {
if affected != current {
return Err(FerrumError::internal(
"cohort maintenance retry must cover the complete current cohort",
));
}
}
}
Ok(Some(retry))
}
pub fn narrower_prefill_tokens(&self, attempted_tokens: usize) -> Option<usize> {
if attempted_tokens <= 1 {
return None;
}
let maximum_next = attempted_tokens
.saturating_sub(attempted_tokens.div_ceil(4))
.max(1);
let proportional = self
.shortfalls()
.iter()
.filter_map(|shortfall| {
let requested = shortfall.requested().get();
let available = shortfall.available().get();
(requested > available).then(|| {
let scaled = (attempted_tokens as u128).saturating_mul(available as u128)
/ requested as u128;
usize::try_from(scaled)
.unwrap_or(usize::MAX)
.clamp(1, attempted_tokens - 1)
})
})
.min();
Some(
proportional
.unwrap_or_else(|| attempted_tokens.div_ceil(2))
.min(maximum_next)
.max(1),
)
}
}
#[derive(Debug, Clone, Serialize)]
pub struct ExecutorRequestStateDeferral {
stage: ExecutorExecutionCapacityStage,
request_ids: Vec<RequestId>,
hazard: crate::vnext::RequestStateHazardDeferral,
}
impl ExecutorRequestStateDeferral {
pub fn new(
stage: ExecutorExecutionCapacityStage,
request_ids: Vec<RequestId>,
hazard: crate::vnext::RequestStateHazardDeferral,
) -> Result<Self> {
let unique = request_ids.iter().collect::<HashSet<_>>();
if request_ids.is_empty() || unique.len() != request_ids.len() {
return Err(FerrumError::internal(
"request-state execution deferral requires a non-empty unique product cohort",
));
}
if hazard.blockers().is_empty() {
return Err(FerrumError::internal(
"request-state execution deferral requires exact blockers",
));
}
Ok(Self {
stage,
request_ids,
hazard,
})
}
pub const fn stage(&self) -> ExecutorExecutionCapacityStage {
self.stage
}
pub fn request_ids(&self) -> &[RequestId] {
&self.request_ids
}
pub const fn hazard(&self) -> &crate::vnext::RequestStateHazardDeferral {
&self.hazard
}
pub fn register_waiter(&self) -> Result<crate::vnext::RequestStateHazardWaitRegistration> {
self.hazard
.register_waiter()
.map_err(|error| FerrumError::backend(error.to_string()))
}
}
#[derive(Debug, Clone, Serialize)]
#[serde(tag = "reason", content = "evidence", rename_all = "snake_case")]
pub enum ExecutorExecutionDeferral {
Capacity(ExecutorExecutionCapacityDeferral),
RequestState(ExecutorRequestStateDeferral),
}
impl ExecutorExecutionDeferral {
pub const fn stage(&self) -> ExecutorExecutionCapacityStage {
match self {
Self::Capacity(deferral) => deferral.stage(),
Self::RequestState(deferral) => deferral.stage(),
}
}
pub const fn as_capacity(&self) -> Option<&ExecutorExecutionCapacityDeferral> {
match self {
Self::Capacity(deferral) => Some(deferral),
Self::RequestState(_) => None,
}
}
pub const fn as_request_state(&self) -> Option<&ExecutorRequestStateDeferral> {
match self {
Self::Capacity(_) => None,
Self::RequestState(deferral) => Some(deferral),
}
}
}
impl From<ExecutorExecutionCapacityDeferral> for ExecutorExecutionDeferral {
fn from(deferral: ExecutorExecutionCapacityDeferral) -> Self {
Self::Capacity(deferral)
}
}
impl From<ExecutorRequestStateDeferral> for ExecutorExecutionDeferral {
fn from(deferral: ExecutorRequestStateDeferral) -> Self {
Self::RequestState(deferral)
}
}
#[cfg(test)]
mod execution_capacity_deferral_tests {
use super::{
ExecutorAdmissionEpochs, ExecutorExecutionCapacityDeferral,
ExecutorExecutionCapacityEvidenceOwner, ExecutorExecutionCapacityStage,
ExecutorExecutionMaintenanceProgress, ExecutorExecutionMaintenanceRetry,
};
use crate::vnext::{
CapacityAvailabilityEpoch, CapacityAvailabilitySource, CapacityWaitCondition,
DeviceCapacityPressure, DeviceCapacityPressureScope, DynamicBackingPressure,
};
use ferrum_types::RequestId;
use std::num::NonZeroU64;
fn test_progress() -> ExecutorExecutionMaintenanceProgress {
ExecutorExecutionMaintenanceProgress {
attempts: 1,
coordinator_id: NonZeroU64::new(19).unwrap(),
mutations: Vec::new(),
latest_capacity_epoch: 5,
}
}
fn test_deferral(stage: ExecutorExecutionCapacityStage) -> ExecutorExecutionCapacityDeferral {
let observed =
CapacityAvailabilityEpoch::new(CapacityAvailabilitySource::ActiveSequenceSlots, 7)
.unwrap();
let condition = CapacityWaitCondition::from_observation(19, vec![observed]).unwrap();
ExecutorExecutionCapacityDeferral::from_backing_pressure(
ExecutorAdmissionEpochs::new(NonZeroU64::new(19).unwrap(), 3, 5),
condition,
test_pressure(),
stage,
)
.unwrap()
}
fn test_pressure() -> DynamicBackingPressure {
DeviceCapacityPressure::new(
DeviceCapacityPressureScope::PlanBudget,
"device.execution-capacity-test".to_owned(),
1,
1,
1,
1,
1,
)
.unwrap()
.into()
}
#[test]
fn prefill_narrowing_is_strict_bounded_and_stops_at_one_token() {
let observed =
CapacityAvailabilityEpoch::new(CapacityAvailabilitySource::ActiveSequenceSlots, 7)
.unwrap();
let condition = CapacityWaitCondition::from_observation(19, vec![observed]).unwrap();
let deferred = ExecutorExecutionCapacityDeferral::from_backing_pressure(
ExecutorAdmissionEpochs::new(NonZeroU64::new(19).unwrap(), 3, 5),
condition,
test_pressure(),
ExecutorExecutionCapacityStage::StepAdmission,
)
.unwrap();
assert_eq!(deferred.narrower_prefill_tokens(342), Some(171));
assert_eq!(deferred.narrower_prefill_tokens(2), Some(1));
assert_eq!(deferred.narrower_prefill_tokens(1), None);
}
#[test]
fn backing_pressure_serializes_one_typed_evidence_owner() {
let deferred = test_deferral(ExecutorExecutionCapacityStage::SequenceExtension);
let serialized = serde_json::to_value(&deferred).unwrap();
assert_eq!(
deferred.evidence().owner(),
ExecutorExecutionCapacityEvidenceOwner::Backing
);
assert!(deferred.shortfalls().is_empty());
assert!(deferred.backing_blockers().is_empty());
assert!(deferred.backing_pressure().is_some());
assert_eq!(serialized["evidence"]["owner"], "backing");
assert_eq!(serialized["evidence"]["kind"], "backing_pressure");
assert!(serialized["evidence"]["pressure"].is_object());
}
#[test]
fn empty_maintenance_receipts_remain_an_ordinary_typed_deferral() {
let request_id = RequestId::new();
let deferred = test_deferral(ExecutorExecutionCapacityStage::SequenceExtension)
.with_relevant_maintenance_retry(2, &[], &[], vec![request_id])
.unwrap();
assert!(deferred.maintenance_retry().is_none());
}
#[test]
fn maintenance_retry_rejects_empty_duplicate_or_unproven_scope() {
let request_id = RequestId::new();
assert!(ExecutorExecutionMaintenanceRetry::new(Vec::new(), test_progress()).is_err());
assert!(ExecutorExecutionMaintenanceRetry::new(
vec![request_id.clone(), request_id.clone()],
test_progress(),
)
.is_err());
assert!(ExecutorExecutionMaintenanceRetry::new(vec![request_id], test_progress()).is_err());
}
#[test]
fn maintenance_retry_scope_is_fail_closed_for_sequence_and_cohort_stages() {
let first = RequestId::new();
let second = RequestId::new();
let retry = |affected_request_ids| ExecutorExecutionMaintenanceRetry {
affected_request_ids,
progress: test_progress(),
};
let mut sequence = test_deferral(ExecutorExecutionCapacityStage::SequenceExtension);
sequence.maintenance_retry = Some(retry(vec![second.clone()]));
assert_eq!(
sequence
.validated_maintenance_retry_scope(&[first.clone(), second.clone()])
.unwrap()
.unwrap()
.affected_request_ids(),
[second.clone()]
);
sequence.maintenance_retry = Some(retry(vec![first.clone(), second.clone()]));
assert!(sequence
.validated_maintenance_retry_scope(&[first.clone(), second.clone()])
.is_err());
let mut cohort = test_deferral(ExecutorExecutionCapacityStage::SubmissionWave);
cohort.maintenance_retry = Some(retry(vec![second.clone()]));
assert!(cohort
.validated_maintenance_retry_scope(&[first.clone(), second.clone()])
.is_err());
cohort.maintenance_retry = Some(retry(vec![first.clone(), second.clone()]));
assert!(cohort
.validated_maintenance_retry_scope(&[first, second])
.unwrap()
.is_some());
}
}
pub enum ExecutorBatchDecodeOutcome {
Completed(Vec<DecodeOutput>),
Deferred(ExecutorExecutionDeferral),
}
pub enum PlanRuntimeBatchDecodeOutcome {
Completed(Vec<PlanRuntimeDecodeOutput>),
Deferred(ExecutorExecutionDeferral),
}
pub struct PlanRuntimePrefillCompletion {
output: PlanRuntimePrefillOutput,
planned_chunk: PrefillChunk,
completed_chunk: PrefillChunk,
capacity_probe_count: u32,
}
impl PlanRuntimePrefillCompletion {
pub fn new(
output: PlanRuntimePrefillOutput,
planned_chunk: PrefillChunk,
completed_chunk: PrefillChunk,
capacity_probe_count: u32,
) -> Result<Self> {
validate_prefill_completion_shape(planned_chunk, completed_chunk, capacity_probe_count)?;
Ok(Self {
output,
planned_chunk,
completed_chunk,
capacity_probe_count,
})
}
pub fn exact(output: PlanRuntimePrefillOutput, chunk: PrefillChunk) -> Self {
Self {
output,
planned_chunk: chunk,
completed_chunk: chunk,
capacity_probe_count: 0,
}
}
pub const fn planned_chunk(&self) -> PrefillChunk {
self.planned_chunk
}
pub const fn completed_chunk(&self) -> PrefillChunk {
self.completed_chunk
}
pub const fn capacity_probe_count(&self) -> u32 {
self.capacity_probe_count
}
pub fn output(&self) -> &PlanRuntimePrefillOutput {
&self.output
}
pub fn validate_for(
&self,
expected_request_id: &RequestId,
expected_planned_chunk: PrefillChunk,
vocabulary_size: usize,
) -> Result<()> {
if self.planned_chunk != expected_planned_chunk {
return Err(FerrumError::backend(format!(
"plan runtime completed prefill frontier {:?}, expected {:?}",
self.planned_chunk.range(),
expected_planned_chunk.range()
)));
}
validate_prefill_completion_shape(
self.planned_chunk,
self.completed_chunk,
self.capacity_probe_count,
)?;
self.output.validate_for_completion(
expected_request_id,
self.completed_chunk,
vocabulary_size,
)
}
pub fn into_parts(self) -> (PlanRuntimePrefillOutput, PrefillChunk, PrefillChunk, u32) {
(
self.output,
self.planned_chunk,
self.completed_chunk,
self.capacity_probe_count,
)
}
}
pub enum PlanRuntimePrefillOutcome {
Completed(PlanRuntimePrefillCompletion),
Deferred(ExecutorExecutionDeferral),
}
pub enum PlanRuntimeBatchPrefillOutcome {
Completed(Vec<PlanRuntimePrefillCompletion>),
NotSubmitted(ExecutorExecutionDeferral),
Unsupported,
}
pub struct ExecutorPrefillCompletion {
output: PrefillOutput,
planned_chunk: PrefillChunk,
completed_chunk: PrefillChunk,
capacity_probe_count: u32,
}
impl ExecutorPrefillCompletion {
pub fn new(
output: PrefillOutput,
planned_chunk: PrefillChunk,
completed_chunk: PrefillChunk,
capacity_probe_count: u32,
) -> Result<Self> {
validate_prefill_completion_shape(planned_chunk, completed_chunk, capacity_probe_count)?;
Ok(Self {
output,
planned_chunk,
completed_chunk,
capacity_probe_count,
})
}
pub fn exact(output: PrefillOutput, chunk: PrefillChunk) -> Self {
Self {
output,
planned_chunk: chunk,
completed_chunk: chunk,
capacity_probe_count: 0,
}
}
pub const fn planned_chunk(&self) -> PrefillChunk {
self.planned_chunk
}
pub const fn completed_chunk(&self) -> PrefillChunk {
self.completed_chunk
}
pub const fn capacity_probe_count(&self) -> u32 {
self.capacity_probe_count
}
pub fn into_parts(self) -> (PrefillOutput, PrefillChunk, PrefillChunk, u32) {
(
self.output,
self.planned_chunk,
self.completed_chunk,
self.capacity_probe_count,
)
}
}
fn validate_prefill_completion_shape(
planned_chunk: PrefillChunk,
completed_chunk: PrefillChunk,
capacity_probe_count: u32,
) -> Result<()> {
if completed_chunk.tokens_processed() != planned_chunk.tokens_processed()
|| completed_chunk.total_prompt_tokens() != planned_chunk.total_prompt_tokens()
|| completed_chunk.tokens_to_process() > planned_chunk.tokens_to_process()
{
return Err(ferrum_types::FerrumError::internal(
"completed prefill chunk is not a non-empty prefix of its planned chunk",
));
}
if completed_chunk != planned_chunk && capacity_probe_count == 0 {
return Err(ferrum_types::FerrumError::internal(
"partial prefill completion requires a failed capacity probe",
));
}
Ok(())
}
pub enum ExecutorPrefillOutcome {
Completed(ExecutorPrefillCompletion),
Deferred(ExecutorExecutionDeferral),
}
pub enum ExecutorBatchPrefillOutcome {
Completed(Vec<ExecutorPrefillCompletion>),
NotSubmitted(ExecutorExecutionDeferral),
Unsupported,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ExecutorPrefillMaintenanceStage {
LogicalCapacity,
PhysicalBacking,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
#[serde(tag = "source", rename_all = "snake_case")]
pub enum ExecutorPrefillMaintenanceBlocker {
Capacity {
domain_id: Option<u32>,
kind: crate::vnext::CapacityShortfallKind,
requested: u64,
available: u64,
current_total: u64,
maximum_total: u64,
},
Backing {
pool_id: String,
domain_id: u32,
lifetime: crate::vnext::DynamicBackingClaimScope,
reason: crate::vnext::DynamicBackingDeferralReason,
requested_bytes: u64,
free_bytes: u64,
largest_contiguous_bytes: u64,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub struct ExecutorPrefillMaintenanceDeferral {
request_id: RequestId,
observed: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
stage: ExecutorPrefillMaintenanceStage,
blockers: Vec<ExecutorPrefillMaintenanceBlocker>,
}
impl ExecutorPrefillMaintenanceDeferral {
pub fn new(
request_id: RequestId,
observed: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
stage: ExecutorPrefillMaintenanceStage,
blockers: Vec<ExecutorPrefillMaintenanceBlocker>,
) -> Result<Self> {
if blockers.is_empty() {
return Err(ferrum_types::FerrumError::request_validation(
"executor prefill maintenance deferral requires at least one blocker",
));
}
if wait_condition.coordinator_id().get() != observed.coordinator_id.get() {
return Err(ferrum_types::FerrumError::request_validation(
"executor prefill maintenance wait condition belongs to a different coordinator",
));
}
Ok(Self {
request_id,
observed,
wait_condition,
stage,
blockers,
})
}
pub fn from_admission(
request_id: &RequestId,
deferred: &crate::vnext::AdmissionDeferred,
) -> Result<Self> {
if deferred.action() != crate::vnext::DeferredAction::AwaitBackingGrowth {
return Err(ferrum_types::FerrumError::internal(
"logical prefill maintenance projection requires AwaitBackingGrowth",
));
}
let blockers = deferred
.blockers()
.iter()
.map(|blocker| ExecutorPrefillMaintenanceBlocker::Capacity {
domain_id: blocker.domain().map(|domain| domain.get()),
kind: blocker.kind(),
requested: blocker.requested().get(),
available: blocker.available().get(),
current_total: blocker.current_total().get(),
maximum_total: blocker.maximum_total().get(),
})
.collect();
Self::new(
request_id.clone(),
ExecutorAdmissionEpochs::from_capacity(deferred.epochs()),
deferred.wait_condition().clone(),
ExecutorPrefillMaintenanceStage::LogicalCapacity,
blockers,
)
}
pub fn from_backing(
request_id: &RequestId,
deferred: &crate::vnext::DynamicBackingDeferred,
) -> Result<Self> {
let blockers = deferred
.blockers()
.iter()
.map(|blocker| ExecutorPrefillMaintenanceBlocker::Backing {
pool_id: blocker.pool_id().as_str().to_string(),
domain_id: blocker.domain_id().get(),
lifetime: deferred.scope(),
reason: blocker.reason(),
requested_bytes: blocker.requested_bytes(),
free_bytes: blocker.free_bytes(),
largest_contiguous_bytes: blocker.largest_contiguous_bytes(),
})
.collect();
Self::new(
request_id.clone(),
ExecutorAdmissionEpochs::from_capacity(deferred.epochs()),
deferred.wait_condition().clone(),
ExecutorPrefillMaintenanceStage::PhysicalBacking,
blockers,
)
}
pub fn request_id(&self) -> &RequestId {
&self.request_id
}
pub const fn observed(&self) -> ExecutorAdmissionEpochs {
self.observed
}
pub fn wait_condition(&self) -> &crate::vnext::CapacityWaitCondition {
&self.wait_condition
}
pub const fn stage(&self) -> ExecutorPrefillMaintenanceStage {
self.stage
}
pub fn blockers(&self) -> &[ExecutorPrefillMaintenanceBlocker] {
&self.blockers
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
#[serde(tag = "outcome", rename_all = "snake_case")]
pub enum ExecutorPrefillMaintenanceOutcome {
NoLongerPending,
RetryAdmission { current: ExecutorAdmissionEpochs },
WaitForRelease {
current: ExecutorAdmissionEpochs,
wait_condition: crate::vnext::CapacityWaitCondition,
pressure: crate::vnext::DynamicBackingPressure,
},
Maintained {
current: ExecutorAdmissionEpochs,
pools_grown: usize,
allocated_bytes: u64,
pools_reclaimed: usize,
chunks_reclaimed: usize,
reclaimed_bytes: u64,
rebalance: Option<crate::vnext::DynamicPoolRebalanceReceipt>,
},
}
#[derive(Debug, Clone)]
pub enum ExecutorPrefillAdmissionDecision {
Admitted(ExecutorPrefillAdmissionReceipt),
Deferred(crate::vnext::AdmissionDeferred),
MaintenanceDeferred(ExecutorPrefillMaintenanceDeferral),
PermanentRejected(crate::vnext::AdmissionRejected),
}
#[async_trait]
pub trait ModelExecutor: Send + Sync {
fn info(&self) -> &ModelInfo;
fn execution_resource_authority(&self) -> ExecutionResourceAuthority {
ExecutionResourceAuthority::LegacyEngine
}
fn admission_limits(&self) -> Result<Option<ExecutorAdmissionLimits>> {
Ok(None)
}
fn resolved_model_plan(&self) -> Option<&crate::vnext::ResolvedModelPlan> {
None
}
fn plan_runtime_resource_snapshot(&self) -> Result<Option<PlanRuntimeResourceSnapshot>> {
Ok(None)
}
fn supports_native_unified_decode(&self) -> bool {
false
}
fn kv_capacity(&self) -> Option<usize> {
None
}
fn attach_execution_event_sink(&self, _sink: Arc<dyn crate::vnext::ExecutionEventSink>) {}
fn execution_capacity_epochs(&self) -> Result<Option<ExecutorAdmissionEpochs>> {
Ok(None)
}
fn write_execution_capacity_snapshot(
&self,
availability: &mut Vec<crate::vnext::CapacityAvailabilityEpoch>,
) -> Result<Option<ExecutorAdmissionEpochs>> {
availability.clear();
self.execution_capacity_epochs()
}
fn register_execution_capacity_waiter(
&self,
_observed: &crate::vnext::CapacityWaitCondition,
) -> Result<Option<ExecutorCapacityWaitRegistration>> {
Ok(None)
}
fn try_admit_prefill(
&self,
_input: ExecutorPrefillAdmission<'_>,
) -> Result<ExecutorPrefillAdmissionDecision> {
Err(ferrum_types::FerrumError::unsupported(
"plan-runtime prefill admission is not implemented",
))
}
fn cancel_prefill_admission(&self, _request_id: &RequestId) -> bool {
false
}
fn write_execution_capacity_release_sources(
&self,
_preemption: &ExecutorExecutionCapacityPreemption,
sources: &mut Vec<crate::vnext::CapacityAvailabilitySource>,
) -> Result<bool> {
sources.clear();
Ok(false)
}
async fn preempt_execution_capacity(
&self,
_preemption: ExecutorExecutionCapacityPreemption,
) -> Result<ExecutorExecutionCapacityPreemptionReceipt> {
Err(FerrumError::unsupported(
"request-scoped execution-capacity preemption is not implemented",
))
}
fn maintain_prefill_backing(
&self,
_request_id: &RequestId,
) -> Result<ExecutorPrefillMaintenanceOutcome> {
Err(ferrum_types::FerrumError::unsupported(
"plan-runtime prefill backing maintenance is not implemented",
))
}
fn reserve_kv_slots(&self, _requests: &[KvSlotRequest]) -> Result<Option<KvSlotReservation>> {
Ok(None)
}
fn kv_slot_capacity_snapshot(&self) -> Option<KvSlotCapacitySnapshot> {
None
}
fn recurrent_state_spec(
&self,
_request_id: &RequestId,
_input_tokens: &[TokenId],
) -> Result<Option<RecurrentStateSpec>> {
Ok(None)
}
async fn prefill(&self, input: &PrefillInput) -> Result<PrefillOutput>;
async fn prefill_with_capacity(&self, input: &PrefillInput) -> Result<ExecutorPrefillOutcome> {
let output = self.prefill(input).await?;
let chunk = match input.chunk {
Some(chunk) => chunk,
None => PrefillChunk::new(0, input.sequence_length(), input.sequence_length())?,
};
Ok(ExecutorPrefillOutcome::Completed(
ExecutorPrefillCompletion::exact(output, chunk),
))
}
async fn batch_prefill(&self, inputs: &[PrefillInput]) -> Result<Vec<PrefillOutput>> {
let mut outputs = Vec::with_capacity(inputs.len());
for input in inputs {
outputs.push(self.prefill(input).await?);
}
Ok(outputs)
}
async fn batch_prefill_with_capacity(
&self,
_inputs: &[PrefillInput],
) -> Result<ExecutorBatchPrefillOutcome> {
Ok(ExecutorBatchPrefillOutcome::Unsupported)
}
async fn plan_runtime_prefill_with_capacity(
&self,
_input: &PlanRuntimePrefillInput,
) -> Result<PlanRuntimePrefillOutcome> {
Err(FerrumError::unsupported(
"tensor-free plan-runtime prefill is not implemented",
))
}
async fn plan_runtime_batch_prefill_with_capacity(
&self,
_inputs: &[PlanRuntimePrefillInput],
) -> Result<PlanRuntimeBatchPrefillOutcome> {
Ok(PlanRuntimeBatchPrefillOutcome::Unsupported)
}
fn discard_plan_runtime_prefill(&self, authority: PlanRuntimePrefillAuthority) -> Result<()> {
self.release_cache(&authority.kv_cache().cache_id());
Ok(())
}
async fn decode(&self, input: &DecodeInput) -> Result<DecodeOutput>;
async fn batch_decode(&self, inputs: &[DecodeInput]) -> Result<Vec<DecodeOutput>> {
let mut outputs = Vec::with_capacity(inputs.len());
for input in inputs {
outputs.push(self.decode(input).await?);
}
Ok(outputs)
}
async fn batch_decode_with_capacity(
&self,
inputs: &[DecodeInput],
) -> Result<ExecutorBatchDecodeOutcome> {
self.batch_decode(inputs)
.await
.map(ExecutorBatchDecodeOutcome::Completed)
}
async fn plan_runtime_batch_decode_with_capacity(
&self,
_inputs: &[PlanRuntimeDecodeInput],
) -> Result<PlanRuntimeBatchDecodeOutcome> {
Err(FerrumError::unsupported(
"tensor-free plan-runtime batch decode is not implemented",
))
}
async fn unified_decode(&self, _batch: &UnifiedBatch) -> Result<Vec<Option<Vec<f32>>>> {
Err(ferrum_types::FerrumError::unsupported(
"unified_decode not implemented for this executor",
))
}
async fn forward(&self, _input: &TensorRef) -> Result<TensorRef> {
Err(ferrum_types::FerrumError::unsupported(
"Full forward pass not supported by this executor",
))
}
async fn truncate_kv(
&self,
_kv_cache: &std::sync::Arc<dyn crate::KvCacheHandle>,
_new_len: usize,
) -> Result<()> {
Ok(())
}
async fn forward_verify(&self, inputs: &[DecodeInput]) -> Result<Vec<DecodeOutput>> {
let mut out = Vec::with_capacity(inputs.len());
for input in inputs {
out.push(self.decode(input).await?);
}
Ok(out)
}
fn capabilities(&self) -> ExecutorCapabilities;
fn status(&self) -> ExecutorStatus;
fn cache_metrics_snapshot(&self) -> Option<serde_json::Value> {
None
}
fn execution_attribution_snapshot(&self) -> Option<serde_json::Value> {
None
}
fn lora_metrics_snapshot(&self) -> Option<serde_json::Value> {
None
}
async fn prepare_startup(&self) -> Result<()> {
Ok(())
}
async fn warmup(&mut self) -> Result<()> {
Ok(())
}
async fn shutdown(&mut self) -> Result<()> {
Ok(())
}
fn complete_cache(&self, completion: ExecutorSequenceCompletion) -> Result<()> {
self.release_cache(completion.cache_id());
Ok(())
}
fn release_cache(&self, _cache_id: &str) {
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExecutorCapabilities {
pub max_batch_size: usize,
pub max_sequence_length: usize,
pub attention_mechanisms: Vec<AttentionType>,
pub supports_dynamic_batching: bool,
pub supports_continuous_batching: bool,
pub supports_speculative_decoding: bool,
pub supports_tensor_parallelism: bool,
pub supports_pipeline_parallelism: bool,
pub supported_dtypes: Vec<ferrum_types::DataType>,
pub supported_devices: Vec<ferrum_types::Device>,
pub memory_requirements: MemoryRequirements,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub enum AttentionType {
MultiHead,
MultiQuery,
GroupedQuery,
Flash,
Paged,
SlidingWindow,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MemoryRequirements {
pub parameter_memory: u64,
pub activation_memory_per_token: usize,
pub kv_cache_memory_per_token: usize,
pub overhead_memory: u64,
}
impl MemoryRequirements {
pub fn calculate_total_memory(
&self,
batch_size: usize,
sequence_length: usize,
num_layers: usize,
) -> u64 {
let activation_mem =
(self.activation_memory_per_token * batch_size * sequence_length) as u64;
let kv_cache_mem =
(self.kv_cache_memory_per_token * batch_size * sequence_length * num_layers) as u64;
self.parameter_memory + activation_mem + kv_cache_mem + self.overhead_memory
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExecutorStatus {
pub state: ExecutorState,
pub is_ready: bool,
pub current_batch_size: usize,
pub prefill_operations: u64,
pub decode_operations: u64,
pub avg_prefill_time_ms: f64,
pub avg_decode_time_ms: f64,
pub memory_usage: ExecutorMemoryUsage,
#[serde(skip)]
pub last_operation: Option<std::time::Instant>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum ExecutorState {
Initializing,
Ready,
Busy,
Error,
Shutdown,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExecutorMemoryUsage {
pub allocated_bytes: usize,
pub used_bytes: usize,
pub peak_bytes: usize,
pub utilization_percent: f32,
}
#[async_trait]
pub trait BatchModelExecutor: ModelExecutor {
async fn batch_prefill(&self, inputs: &[PrefillInput]) -> Result<Vec<PrefillOutput>>;
async fn batch_decode(&self, inputs: &[DecodeInput]) -> Result<Vec<DecodeOutput>>;
fn optimal_batch_size(&self) -> usize;
fn supports_batch_size(&self, batch_size: usize) -> bool;
}
#[async_trait]
pub trait SpeculativeExecutor: ModelExecutor {
async fn speculative_decode(
&self,
input: &DecodeInput,
draft_tokens: &[ferrum_types::TokenId],
acceptance_threshold: f32,
) -> Result<SpeculativeDecodeOutput>;
}
#[derive(Debug, Clone)]
pub struct SpeculativeDecodeOutput {
pub accepted_tokens: Vec<ferrum_types::TokenId>,
pub next_logits: TensorRef,
pub kv_cache: Arc<dyn KvCacheHandle>,
pub acceptance_count: usize,
}
#[async_trait]
pub trait ModelExecutorFactory: Send + Sync {
async fn create_executor(&self, config: &ExecutorConfig) -> Result<Box<dyn ModelExecutor>>;
async fn create_batch_executor(
&self,
config: &ExecutorConfig,
) -> Result<Box<dyn BatchModelExecutor>>;
fn supported_types(&self) -> Vec<ExecutorType>;
fn validate_config(&self, config: &ExecutorConfig) -> Result<()>;
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExecutorConfig {
pub model_info: ModelInfo,
pub device: ferrum_types::Device,
pub dtype: ferrum_types::DataType,
pub max_batch_size: usize,
pub max_sequence_length: usize,
pub attention_config: ExecutorAttentionConfig,
pub memory_config: ExecutorMemoryConfig,
pub optimization_config: OptimizationConfig,
pub executor_options: HashMap<String, serde_json::Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExecutorAttentionConfig {
pub attention_type: AttentionType,
pub enable_flash_attention: bool,
pub enable_paged_attention: bool,
pub block_size: Option<usize>,
pub sliding_window_size: Option<usize>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExecutorMemoryConfig {
pub enable_memory_pooling: bool,
pub memory_pool_size: Option<usize>,
pub enable_kv_cache_sharing: bool,
pub max_memory_usage: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OptimizationConfig {
pub enable_cuda_graphs: bool,
pub enable_kernel_fusion: bool,
pub enable_mixed_precision: bool,
pub optimization_level: u8,
pub custom_flags: HashMap<String, bool>,
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub enum ExecutorType {
Sequential,
Batch,
ContinuousBatch,
Speculative,
PipelineParallel,
TensorParallel,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ExecutorMetrics {
pub total_operations: u64,
pub prefill_operations: u64,
pub decode_operations: u64,
pub avg_prefill_latency: f64,
pub avg_decode_latency: f64,
pub p95_prefill_latency: f64,
pub p95_decode_latency: f64,
pub throughput_tps: f64,
pub memory_efficiency: f32,
pub batch_utilization: f32,
}
pub trait ExecutorRegistry: Send + Sync {
fn register(&mut self, name: &str, executor: Box<dyn ModelExecutor>) -> Result<()>;
fn get(&self, name: &str) -> Option<&dyn ModelExecutor>;
fn remove(&mut self, name: &str) -> Option<Box<dyn ModelExecutor>>;
fn list_names(&self) -> Vec<String>;
fn get_metrics(&self, name: &str) -> Option<ExecutorMetrics>;
}