use crate::{
cache::{
LayerCachePolicy, StateTensorDimension, StateTensorDtype, StateTensorPolicy,
StateTensorPresence, StateTensorRole,
},
AttentionPolicy, LayerSchedule, ObservationKind, Observed,
};
use serde::{Deserialize, Serialize};
use std::num::NonZeroU8;
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
pub struct InputModalities {
pub text: bool,
pub image: bool,
pub audio: bool,
pub video: bool,
}
impl InputModalities {
pub const TEXT: Self = Self {
text: true,
image: false,
audio: false,
video: false,
};
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "strategy", rename_all = "snake_case")]
pub enum CacheStateStrategy {
FullKv,
SlidingKv {
window: u64,
},
SlidingKey {
window: u64,
layers: u64,
pooling_layers: u64,
},
MixedKv {
full_layers: u64,
sliding: Vec<SlidingWindowLayerCount>,
},
SharedFullKv {
cached_layers: u64,
shared_layers: u64,
full_attention_layers: u64,
sliding_attention: Vec<SlidingWindowLayerCount>,
},
CompressedMla {
latent_width: u64,
rotary_width: u64,
},
HybridRecurrent {
full_attention_layers: u64,
sliding_attention: Vec<SlidingWindowLayerCount>,
recurrent_layers: u64,
},
Multimodal {
decoder: Box<CacheStateStrategy>,
media_consumes_decoder_positions: bool,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub struct SlidingWindowLayerCount {
pub window: u64,
pub layers: u64,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum EstimationCompleteness {
Complete,
Conservative,
PersistentStateOnly,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ModelCapabilities {
pub effective_model_type: String,
pub native_max_context: Observed<u64>,
pub effective_max_context: Observed<u64>,
pub state_strategy: CacheStateStrategy,
pub modalities: InputModalities,
pub estimation: EstimationCompleteness,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
pub struct InputTokenCount {
pub text_tokens: u64,
pub media_positions: u64,
pub model_positions: u64,
pub kind: ObservationKind,
media_execution_workspace_bytes: u64,
media_execution_workspace_kind: ObservationKind,
}
impl InputTokenCount {
pub const fn text(tokens: u64) -> Self {
Self {
text_tokens: tokens,
media_positions: 0,
model_positions: tokens,
kind: ObservationKind::Exact,
media_execution_workspace_bytes: 0,
media_execution_workspace_kind: ObservationKind::Exact,
}
}
pub const fn prepared(
text_tokens: u64,
media_positions: u64,
model_positions: u64,
media_execution_workspace_bytes: u64,
media_execution_workspace_kind: ObservationKind,
) -> Self {
Self {
text_tokens,
media_positions,
model_positions,
kind: ObservationKind::Exact,
media_execution_workspace_bytes,
media_execution_workspace_kind,
}
}
pub const fn media_execution_workspace_bytes(&self) -> u64 {
self.media_execution_workspace_bytes
}
pub const fn media_execution_workspace_kind(&self) -> ObservationKind {
self.media_execution_workspace_kind
}
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct StateMemoryAssumptions {
pub floating_state_dtype_bytes: NonZeroU8,
pub batch_size: u64,
pub requested_positions: u64,
pub sliding_window_bounds: Vec<u64>,
pub allocation_granularity: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct RuntimeStateEstimate {
pub fixed_state_bytes: u64,
pub bytes_per_position_per_batch: u64,
pub context_state_bytes: u64,
pub multimodal_embedding_bytes: u64,
pub media_execution_workspace_bytes: u64,
pub requested_state_bytes: u64,
pub assumptions: StateMemoryAssumptions,
pub completeness: EstimationCompleteness,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum PhysicalMemorySemantics {
Unified,
SeparateTiers,
Unknown,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct StaticMemoryReport {
pub logical_parameter_bytes: Observed<u64>,
pub current_host_resident_bytes: Observed<u64>,
pub current_device_resident_bytes: Observed<u64>,
pub planned_disk_backed_bytes: Observed<u64>,
pub backend_active_allocation_bytes: Observed<u64>,
pub backend_allocator_cache_bytes: Observed<u64>,
pub physical_semantics: PhysicalMemorySemantics,
pub currently_cached_shards: Observed<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct AvailableMemory {
pub physical_memory_bytes: Observed<u64>,
pub available_memory_bytes: Observed<u64>,
pub physical_semantics: PhysicalMemorySemantics,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Serialize, Deserialize)]
pub struct AdmissionRequest {
pub input: InputTokenCount,
pub max_output_tokens: u64,
pub batch_size: u64,
pub safety_reserve_bytes: u64,
pub application_memory_budget_bytes: Option<u64>,
pub require_complete_estimate: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Admission {
pub requested_positions: u64,
pub state: RuntimeStateEstimate,
pub incremental_required_bytes: u64,
pub available_memory_bytes: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum AdmissionRejection {
PromptExceedsContext {
prompt_positions: u64,
maximum_positions: u64,
},
OutputHeadroomExceedsContext {
prompt_positions: u64,
output_tokens: u64,
maximum_positions: u64,
},
MemoryBudgetExceeded {
required_bytes: u64,
budget_bytes: u64,
},
InsufficientAvailableMemory {
required_bytes: u64,
available_bytes: u64,
},
AvailableMemoryUnavailable {
reason: String,
},
EstimationUnsupported {
reason: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "status", rename_all = "snake_case")]
pub enum AdmissionResult {
Admitted(Admission),
Rejected(AdmissionRejection),
}
#[derive(Debug, thiserror::Error, Clone, PartialEq, Eq)]
pub enum CapabilityError {
#[error("invalid model capability field {field}: {detail}")]
InvalidConfiguration {
field: &'static str,
detail: String,
},
#[error("capability arithmetic overflow while computing {operation}")]
ArithmeticOverflow {
operation: &'static str,
},
#[error("unsupported prepared input for {architecture}: {reason}")]
UnsupportedInput {
architecture: String,
reason: String,
},
#[error("capability observation failed: {0}")]
Observation(String),
}
#[derive(Debug, Clone, Eq, PartialEq, Serialize, Deserialize)]
pub struct StateMemoryLayout {
layer_layout: LayerSchedule<LayerCachePolicy>,
layer_prefix_offsets: Vec<i32>,
pub hidden_size: u64,
pub allocation_granularity: u64,
pub completeness: EstimationCompleteness,
}
impl StateMemoryLayout {
pub fn new(
layer_layout: LayerSchedule<LayerCachePolicy>,
layer_prefix_offsets: Vec<i32>,
hidden_size: u64,
allocation_granularity: u64,
completeness: EstimationCompleteness,
) -> Result<Self, CapabilityError> {
if layer_layout.is_empty()
|| layer_prefix_offsets.len() != layer_layout.len()
|| layer_prefix_offsets.iter().any(|offset| *offset > 0)
|| hidden_size == 0
|| allocation_granularity == 0
{
let (field, detail) = if layer_layout.is_empty() {
(
"layer_layout",
"must contain at least one executable state layer",
)
} else if layer_prefix_offsets.len() != layer_layout.len() {
(
"layer_prefix_offsets",
"must contain one entry per executable state layer",
)
} else if layer_prefix_offsets.iter().any(|offset| *offset > 0) {
(
"layer_prefix_offsets",
"must not advance beyond the request token frontier",
)
} else if hidden_size == 0 {
("hidden_size", "must be positive")
} else {
("allocation_granularity", "must be positive")
};
return Err(CapabilityError::InvalidConfiguration {
field,
detail: detail.into(),
});
}
for (layer, policy) in layer_layout.iter().enumerate() {
policy
.validate()
.map_err(|error| CapabilityError::InvalidConfiguration {
field: "layer_layout",
detail: format!("invalid state policy at layer {layer}: {error}"),
})?;
}
Ok(Self {
layer_layout,
layer_prefix_offsets,
hidden_size,
allocation_granularity,
completeness,
})
}
pub const fn layer_layout(&self) -> &LayerSchedule<LayerCachePolicy> {
&self.layer_layout
}
pub fn layer_prefix_offsets(&self) -> &[i32] {
&self.layer_prefix_offsets
}
}
fn checked_add(left: u64, right: u64, operation: &'static str) -> Result<u64, CapabilityError> {
left.checked_add(right)
.ok_or(CapabilityError::ArithmeticOverflow { operation })
}
fn checked_mul(left: u64, right: u64, operation: &'static str) -> Result<u64, CapabilityError> {
left.checked_mul(right)
.ok_or(CapabilityError::ArithmeticOverflow { operation })
}
fn attention_scalars_per_position(policy: &LayerCachePolicy) -> Result<u64, CapabilityError> {
let scalars = match policy {
LayerCachePolicy::KeyValue {
num_key_value_heads,
head_dim,
..
}
| LayerCachePolicy::KeyValueWithFixedState {
num_key_value_heads,
head_dim,
..
} => checked_mul(
checked_mul(
u64::from(num_key_value_heads.get()),
u64::from(head_dim.get()),
"key/value heads times head dimension",
)?,
2,
"key plus value scalars",
)?,
LayerCachePolicy::KeyOnly {
num_key_heads,
head_dim,
..
}
| LayerCachePolicy::KeyOnlyWithFixedState {
num_key_heads,
head_dim,
..
} => checked_mul(
u64::from(num_key_heads.get()),
u64::from(head_dim.get()),
"key heads times head dimension",
)?,
LayerCachePolicy::CompressedLatentRotary {
latent_dim,
rotary_dim,
..
} => checked_add(
u64::from(latent_dim.get()),
u64::from(rotary_dim.get()),
"compressed latent plus rotary width",
)?,
LayerCachePolicy::NoState | LayerCachePolicy::FixedState { .. } => 0,
};
Ok(scalars)
}
fn is_context_dependent_dimension(dimension: &StateTensorDimension) -> bool {
matches!(
dimension,
StateTensorDimension::PrefixTokens
| StateTensorDimension::PrefixTokensDiv(_)
| StateTensorDimension::PrefixTokensRem(_)
)
}
fn state_tensor_dtype_bytes(tensor: &StateTensorPolicy, floating_scalar_bytes: u64) -> u64 {
match tensor.dtype {
StateTensorDtype::Floating => floating_scalar_bytes,
StateTensorDtype::Float32 | StateTensorDtype::Int32 | StateTensorDtype::Uint32 => 4,
}
}
fn state_tensor_is_present(tensor: &StateTensorPolicy, prefix_tokens: usize) -> bool {
match tensor.presence {
StateTensorPresence::Required => true,
StateTensorPresence::Optional => !matches!(tensor.role, StateTensorRole::PrefixEmbedding),
StateTensorPresence::PrefixRemainderNonZero(divisor) => {
!prefix_tokens.is_multiple_of(divisor.get() as usize)
}
StateTensorPresence::PrefixAtLeast(divisor) => prefix_tokens >= divisor.get() as usize,
}
}
fn state_tensor_bytes(
tensor: &StateTensorPolicy,
batch_size: usize,
prefix_tokens: usize,
floating_scalar_bytes: u64,
) -> Result<u64, CapabilityError> {
if !state_tensor_is_present(tensor, prefix_tokens) {
return Ok(0);
}
let shape = tensor
.resolved_shape(batch_size, prefix_tokens)
.map_err(|error| CapabilityError::InvalidConfiguration {
field: "layer_layout",
detail: error.to_string(),
})?;
let scalars = shape.into_iter().try_fold(1_u64, |scalars, dimension| {
checked_mul(
scalars,
u64::try_from(dimension).map_err(|_| CapabilityError::InvalidConfiguration {
field: "layer_layout",
detail: "runtime state tensor has a negative resolved dimension".into(),
})?,
"runtime state tensor scalar count",
)
})?;
checked_mul(
scalars,
state_tensor_dtype_bytes(tensor, floating_scalar_bytes),
"runtime state tensor bytes",
)
}
fn state_tensor_bytes_per_position_per_batch(
tensor: &StateTensorPolicy,
floating_scalar_bytes: u64,
) -> Result<u64, CapabilityError> {
let mut scalars = 1_u64;
let mut divisor = 1_u64;
let mut unbounded = false;
for dimension in &tensor.shape {
match dimension {
StateTensorDimension::Batch | StateTensorDimension::Scalar => {}
StateTensorDimension::Fixed(value) => {
scalars =
checked_mul(scalars, u64::from(value.get()), "state growth scalar count")?;
}
StateTensorDimension::PrefixTokens => unbounded = true,
StateTensorDimension::PrefixTokensDiv(value) => {
unbounded = true;
divisor = checked_mul(divisor, u64::from(value.get()), "state growth divisor")?;
}
StateTensorDimension::PrefixTokensRem(_) => return Ok(0),
}
}
if !unbounded {
return Ok(0);
}
let bytes = checked_mul(
scalars,
state_tensor_dtype_bytes(tensor, floating_scalar_bytes),
"state growth bytes",
)?;
Ok(bytes.div_ceil(divisor))
}
pub fn estimate_runtime_state(
layout: &StateMemoryLayout,
input: InputTokenCount,
max_output_tokens: u64,
batch_size: u64,
floating_state_dtype_bytes: NonZeroU8,
) -> Result<RuntimeStateEstimate, CapabilityError> {
if batch_size == 0 {
return Err(CapabilityError::InvalidConfiguration {
field: "batch_size",
detail: "must be positive".into(),
});
}
let requested_positions = checked_add(
input.model_positions,
max_output_tokens,
"prompt plus output positions",
)?;
let floating_scalar_bytes = u64::from(floating_state_dtype_bytes.get());
let batch_size_usize =
usize::try_from(batch_size).map_err(|_| CapabilityError::InvalidConfiguration {
field: "batch_size",
detail: "exceeds the runtime state shape range".into(),
})?;
let mut fixed_state_bytes = 0;
let mut context_state_bytes = 0;
let mut unbounded_per_position = 0;
let mut sliding_window_bounds = Vec::new();
for (layer, policy) in layout.layer_layout.iter().enumerate() {
let layer_positions = requested_positions
.saturating_sub(u64::from(layout.layer_prefix_offsets[layer].unsigned_abs()));
let layer_positions_usize = usize::try_from(layer_positions).map_err(|_| {
CapabilityError::InvalidConfiguration {
field: "requested_positions",
detail: "exceeds the runtime state shape range".into(),
}
})?;
if let Some(attention) = policy.attention() {
let per_position = attention_scalars_per_position(policy)?;
let retained = match attention {
AttentionPolicy::Sliding { window } => {
let window = u64::from(window.get());
sliding_window_bounds.push(window);
layer_positions.min(window)
}
AttentionPolicy::Full => {
let adjustment = layout.allocation_granularity - 1;
checked_add(layer_positions, adjustment, "cache allocation rounding")?
/ layout.allocation_granularity
* layout.allocation_granularity
}
};
let bytes = checked_mul(
checked_mul(
checked_mul(per_position, retained, "attention context scalars")?,
batch_size,
"attention context batch",
)?,
floating_scalar_bytes,
"attention context bytes",
)?;
context_state_bytes =
checked_add(context_state_bytes, bytes, "context state byte total")?;
if matches!(attention, AttentionPolicy::Full) {
unbounded_per_position = checked_add(
unbounded_per_position,
checked_mul(
per_position,
floating_scalar_bytes,
"unbounded bytes per position",
)?,
"unbounded bytes-per-position total",
)?;
}
}
for tensor in policy.fixed_state() {
let bytes = state_tensor_bytes(
tensor,
batch_size_usize,
layer_positions_usize,
floating_scalar_bytes,
)?;
if tensor.shape.iter().any(is_context_dependent_dimension) {
context_state_bytes =
checked_add(context_state_bytes, bytes, "context state byte total")?;
unbounded_per_position = checked_add(
unbounded_per_position,
state_tensor_bytes_per_position_per_batch(tensor, floating_scalar_bytes)?,
"unbounded bytes-per-position total",
)?;
} else {
fixed_state_bytes =
checked_add(fixed_state_bytes, bytes, "fixed state byte total")?;
}
}
}
sliding_window_bounds.sort_unstable();
sliding_window_bounds.dedup();
let multimodal_embedding_bytes = checked_mul(
checked_mul(
checked_mul(
input.media_positions,
layout.hidden_size,
"media positions times hidden size",
)?,
batch_size,
"media embeddings times batch",
)?,
floating_scalar_bytes,
"media embedding bytes",
)?;
let media_execution_workspace_bytes = checked_mul(
input.media_execution_workspace_bytes,
batch_size,
"media execution workspace times batch",
)?;
let requested_state_bytes = checked_add(
checked_add(
checked_add(
fixed_state_bytes,
context_state_bytes,
"fixed plus context state",
)?,
multimodal_embedding_bytes,
"persistent plus multimodal embedding state",
)?,
media_execution_workspace_bytes,
"persistent plus media execution workspace",
)?;
let completeness = if input.media_positions == 0
|| input.media_execution_workspace_kind == ObservationKind::Exact
{
layout.completeness
} else {
EstimationCompleteness::Conservative
};
Ok(RuntimeStateEstimate {
fixed_state_bytes,
bytes_per_position_per_batch: unbounded_per_position,
context_state_bytes,
multimodal_embedding_bytes,
media_execution_workspace_bytes,
requested_state_bytes,
assumptions: StateMemoryAssumptions {
floating_state_dtype_bytes,
batch_size,
requested_positions,
sliding_window_bounds,
allocation_granularity: layout.allocation_granularity,
},
completeness,
})
}
pub fn apply_admission_policy(
capabilities: &ModelCapabilities,
request: AdmissionRequest,
state: RuntimeStateEstimate,
available: Option<&AvailableMemory>,
) -> Result<AdmissionResult, CapabilityError> {
let maximum = match &capabilities.effective_max_context {
Observed::Available { value, .. } => *value,
Observed::Unsupported { reason } | Observed::Unavailable { reason } => {
return Ok(AdmissionResult::Rejected(
AdmissionRejection::EstimationUnsupported {
reason: reason.clone(),
},
));
}
};
if request.input.model_positions > maximum {
return Ok(AdmissionResult::Rejected(
AdmissionRejection::PromptExceedsContext {
prompt_positions: request.input.model_positions,
maximum_positions: maximum,
},
));
}
let requested_positions = checked_add(
request.input.model_positions,
request.max_output_tokens,
"admission prompt plus output",
)?;
if requested_positions > maximum {
return Ok(AdmissionResult::Rejected(
AdmissionRejection::OutputHeadroomExceedsContext {
prompt_positions: request.input.model_positions,
output_tokens: request.max_output_tokens,
maximum_positions: maximum,
},
));
}
if request.require_complete_estimate
&& state.completeness == EstimationCompleteness::PersistentStateOnly
{
return Ok(AdmissionResult::Rejected(
AdmissionRejection::EstimationUnsupported {
reason: format!(
"architecture estimator coverage is {:?}",
state.completeness
),
},
));
}
let incremental_required_bytes = checked_add(
state.requested_state_bytes,
request.safety_reserve_bytes,
"state plus safety reserve",
)?;
if let Some(budget_bytes) = request.application_memory_budget_bytes {
if incremental_required_bytes > budget_bytes {
return Ok(AdmissionResult::Rejected(
AdmissionRejection::MemoryBudgetExceeded {
required_bytes: incremental_required_bytes,
budget_bytes,
},
));
}
}
let available_memory_bytes = match available {
Some(report) => match &report.available_memory_bytes {
Observed::Available { value, .. } => Some(*value),
Observed::Unsupported { reason } | Observed::Unavailable { reason } => {
return Ok(AdmissionResult::Rejected(
AdmissionRejection::AvailableMemoryUnavailable {
reason: reason.clone(),
},
))
}
},
None => None,
};
if let Some(available_bytes) = available_memory_bytes {
if incremental_required_bytes > available_bytes {
return Ok(AdmissionResult::Rejected(
AdmissionRejection::InsufficientAvailableMemory {
required_bytes: incremental_required_bytes,
available_bytes,
},
));
}
}
Ok(AdmissionResult::Admitted(Admission {
requested_positions,
state,
incremental_required_bytes,
available_memory_bytes,
}))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn state_estimation_and_admission_are_backend_independent() {
let policies = (0..2)
.map(|_| LayerCachePolicy::key_only(AttentionPolicy::Full, 1, 8).unwrap())
.collect::<Vec<_>>();
let layout = StateMemoryLayout::new(
LayerSchedule::new(2, policies).unwrap(),
vec![0; 2],
32,
8,
EstimationCompleteness::Complete,
)
.unwrap();
let input = InputTokenCount::text(5);
let state =
estimate_runtime_state(&layout, input, 2, 1, NonZeroU8::new(4).unwrap()).unwrap();
assert_eq!(state.assumptions.requested_positions, 7);
assert_eq!(state.context_state_bytes, 512);
let capabilities = ModelCapabilities {
effective_model_type: "mock".into(),
native_max_context: Observed::exact(16, "mock"),
effective_max_context: Observed::exact(16, "mock"),
state_strategy: CacheStateStrategy::FullKv,
modalities: InputModalities::TEXT,
estimation: EstimationCompleteness::Complete,
};
let serialized = serde_json::to_value(&capabilities).unwrap();
assert_eq!(serialized["effective_model_type"], "mock");
assert!(serialized.get("model_type").is_none());
assert!(matches!(
apply_admission_policy(
&capabilities,
AdmissionRequest {
input,
max_output_tokens: 2,
batch_size: 1,
safety_reserve_bytes: 0,
application_memory_budget_bytes: Some(1024),
require_complete_estimate: true
},
state,
None
)
.unwrap(),
AdmissionResult::Admitted(_)
));
}
#[test]
fn admission_rejections_are_portable_and_fail_closed() {
let capabilities = ModelCapabilities {
effective_model_type: "mock".into(),
native_max_context: Observed::exact(8, "mock"),
effective_max_context: Observed::exact(8, "mock"),
state_strategy: CacheStateStrategy::FullKv,
modalities: InputModalities::TEXT,
estimation: EstimationCompleteness::Complete,
};
let state = RuntimeStateEstimate {
fixed_state_bytes: 0,
bytes_per_position_per_batch: 0,
context_state_bytes: 0,
multimodal_embedding_bytes: 0,
media_execution_workspace_bytes: 0,
requested_state_bytes: 0,
assumptions: StateMemoryAssumptions {
floating_state_dtype_bytes: NonZeroU8::new(4).unwrap(),
batch_size: 1,
requested_positions: 9,
sliding_window_bounds: Vec::new(),
allocation_granularity: 1,
},
completeness: EstimationCompleteness::Complete,
};
let request = AdmissionRequest {
input: InputTokenCount::text(7),
max_output_tokens: 2,
batch_size: 1,
safety_reserve_bytes: 0,
application_memory_budget_bytes: None,
require_complete_estimate: true,
};
assert!(matches!(
apply_admission_policy(&capabilities, request, state, None).unwrap(),
AdmissionResult::Rejected(AdmissionRejection::OutputHeadroomExceedsContext { .. })
));
let unavailable = AvailableMemory {
physical_memory_bytes: Observed::unavailable("not reported"),
available_memory_bytes: Observed::unavailable("not reported"),
physical_semantics: PhysicalMemorySemantics::Unknown,
};
let request = AdmissionRequest {
input: InputTokenCount::text(1),
max_output_tokens: 0,
batch_size: 1,
safety_reserve_bytes: 0,
application_memory_budget_bytes: None,
require_complete_estimate: true,
};
let state = estimate_runtime_state(
&StateMemoryLayout::new(
LayerSchedule::new(1, vec![LayerCachePolicy::NoState]).unwrap(),
vec![0],
1,
1,
EstimationCompleteness::Complete,
)
.unwrap(),
request.input,
0,
1,
NonZeroU8::new(4).unwrap(),
)
.unwrap();
assert!(matches!(
apply_admission_policy(&capabilities, request, state, Some(&unavailable)).unwrap(),
AdmissionResult::Rejected(AdmissionRejection::AvailableMemoryUnavailable { .. })
));
}
#[test]
fn capability_and_memory_schemas_round_trip_without_a_backend() {
let report = StaticMemoryReport {
logical_parameter_bytes: Observed::exact(1_024, "mock catalog"),
current_host_resident_bytes: Observed::exact(512, "mock ledger"),
current_device_resident_bytes: Observed::exact(512, "mock ledger"),
planned_disk_backed_bytes: Observed::exact(0, "mock plan"),
backend_active_allocation_bytes: Observed::unavailable("no allocator probe"),
backend_allocator_cache_bytes: Observed::unsupported("no allocator cache"),
physical_semantics: PhysicalMemorySemantics::SeparateTiers,
currently_cached_shards: Observed::exact(1, "mock store"),
};
let encoded = serde_json::to_string(&report).unwrap();
let decoded: StaticMemoryReport = serde_json::from_str(&encoded).unwrap();
assert_eq!(decoded, report);
}
}