use std::{
collections::{BTreeMap, BTreeSet},
sync::{Arc, Mutex},
time::Instant,
};
use anyhow::{Context, Result, bail};
use skippy_protocol::{FlashAttentionType, LoadMode, SplitMode, StageConfig};
use skippy_runtime::{
ActivationBoundaryDesc, ActivationFrame, DecodeBatchRequest, DecodeFrameBatchOutput,
DecodeFrameBatchRequest, FlashAttentionType as RuntimeFlashAttentionType,
GenerationSignalWindow, GlmDsaPolicy as RuntimeGlmDsaPolicy, IterationBatchOutput,
IterationBatchPhase, IterationBatchRequest, MediaInput, MediaPrefill, MediaPrefillFrame,
ModelStateKind, ModelWorkload, MtpSource, NativeMtpDraft, RuntimeConfig, RuntimeKvPage,
RuntimeKvPageDesc, RuntimeLoadMode, SamplingConfig, SpeechAudio, SpeechSynthesisConfig,
SplitMode as RuntimeSplitMode, StageModel, StageSession, TokenSignal, WorkloadInfo,
parse_cache_type,
};
mod frame_operations;
mod lane_lifecycle;
pub mod lifecycle;
mod restore_transaction;
mod state_transfer;
pub use lifecycle::{SessionLifecycleEvent, SessionLifecycleObserver};
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RuntimeLaunchOverrides {
pub n_threads: Option<usize>,
pub n_threads_batch: Option<usize>,
pub mtp_source: MtpSource,
}
pub struct RuntimeState {
pub model: StageModel,
compute_meter: Arc<crate::compute_meter::StageComputeMeter>,
layer_start: u32,
layer_end: u32,
lane_count: u32,
ctx_size: u32,
next_lane_index: usize,
free_lane_indices: Vec<usize>,
sessions: BTreeMap<String, RuntimeLaneSession>,
idle_sessions: Vec<RuntimeLaneSession>,
max_idle_sessions: Option<usize>,
session_token_counts: BTreeMap<String, u64>,
session_resident_prefixes: BTreeMap<String, ResidentLanePrefix>,
session_lifecycle_observer: Option<Arc<dyn SessionLifecycleObserver>>,
#[cfg(test)]
modelless_for_test: bool,
}
struct RuntimeLaneSession {
index: usize,
session: StageSession,
resident_prefix: Option<ResidentLanePrefix>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RuntimeSessionLaneStats {
pub index: usize,
pub active: bool,
pub session_id: Option<String>,
pub token_count: Option<u64>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RuntimeSessionStats {
pub lane_count: usize,
pub active_sessions: usize,
pub idle_sessions: usize,
pub idle_resident_prefixes: usize,
pub tracked_token_counts: usize,
pub max_session_tokens: u64,
pub total_session_tokens: u64,
pub graphs_reused: u64,
pub tokens_evaluated: u64,
pub lanes: Vec<RuntimeSessionLaneStats>,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct RuntimeSessionDropStats {
pub reset_session: bool,
pub reset_ms: f64,
pub preserved_resident_prefix: bool,
pub lane_discarded: bool,
pub lane_discard_reason: Option<String>,
pub stats_after: RuntimeSessionStats,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct RuntimeSessionAlignStats {
pub before_token_count: u64,
pub after_token_count: u64,
}
pub struct RuntimeDecodeBatchRequest<'a> {
pub session_id: &'a str,
pub token_id: i32,
pub sampling: Option<&'a SamplingConfig>,
}
pub struct RuntimeDecodeFrameBatchRequest<'a> {
pub session_id: &'a str,
pub token_id: i32,
pub sampling: Option<&'a SamplingConfig>,
pub input: Option<&'a ActivationFrame>,
}
pub struct RuntimeIterationBatchRequest<'a> {
pub session_id: &'a str,
pub token_ids: &'a [i32],
pub positions: &'a [i32],
pub sampling: Option<&'a SamplingConfig>,
pub input: Option<&'a ActivationFrame>,
pub sample_last: bool,
pub phase: IterationBatchPhase,
}
#[derive(Debug, Clone)]
struct ResidentLanePrefix {
page_id: String,
token_count: u64,
}
impl RuntimeState {
pub fn workload_info(&self) -> Result<WorkloadInfo> {
self.model.workload_info()
}
pub fn embed(
&mut self,
session_id: &str,
token_ids: &[i32],
dimensions: usize,
) -> Result<Vec<f32>> {
let embedding = self.session(session_id)?.embed(token_ids, dimensions)?;
self.session_token_counts.insert(
session_id.to_string(),
u64::try_from(token_ids.len()).context("embedding token count exceeds u64")?,
);
Ok(embedding)
}
pub fn rerank(
&mut self,
session_id: &str,
query: &str,
document: &str,
) -> Result<(f32, usize)> {
let result = self.session(session_id)?.rerank(query, document)?;
self.session_token_counts.insert(
session_id.to_string(),
u64::try_from(result.1).context("rerank token count exceeds u64")?,
);
Ok(result)
}
pub fn encode_prompt(&mut self, session_id: &str, token_ids: &[i32]) -> Result<i32> {
let decoder_start = self.session(session_id)?.encode_prompt(token_ids)?;
self.session_token_counts.insert(session_id.to_string(), 0);
Ok(decoder_start)
}
pub fn input_activation_boundary(&self) -> Option<ActivationBoundaryDesc> {
self.model.input_activation_boundary()
}
pub fn output_activation_boundary(&self) -> Option<ActivationBoundaryDesc> {
self.model.output_activation_boundary()
}
#[cfg(test)]
pub(crate) fn new_modelless_for_test(lane_count: u32) -> Self {
Self::new_modelless_with_capacity_for_test(lane_count, 0)
}
#[cfg(test)]
pub(crate) fn new_modelless_with_capacity_for_test(lane_count: u32, ctx_size: u32) -> Self {
Self {
model: StageModel::new_dummy(),
layer_start: 0,
layer_end: 1,
lane_count,
ctx_size,
next_lane_index: 0,
free_lane_indices: Vec::new(),
sessions: BTreeMap::new(),
idle_sessions: Vec::new(),
max_idle_sessions: None,
session_token_counts: BTreeMap::new(),
session_resident_prefixes: BTreeMap::new(),
session_lifecycle_observer: None,
compute_meter: Arc::default(),
modelless_for_test: true,
}
}
#[must_use]
pub fn with_session_lifecycle_observer(
mut self,
observer: Arc<dyn SessionLifecycleObserver>,
) -> Self {
self.session_lifecycle_observer = Some(observer);
self
}
pub(crate) fn notify_session_lifecycle(&self, event: SessionLifecycleEvent) {
if let Some(observer) = self.session_lifecycle_observer.as_ref() {
observer.observe(event);
}
}
#[cfg(test)]
pub(crate) fn track_session_tokens_for_test(&mut self, session_id: &str, token_count: u64) {
self.session_token_counts
.insert(session_id.to_string(), token_count);
}
pub fn lane_count(&self) -> u32 {
self.lane_count
}
pub fn compute_meter(&self) -> Arc<crate::compute_meter::StageComputeMeter> {
self.compute_meter.clone()
}
pub fn set_compute_meter(&mut self, meter: Arc<crate::compute_meter::StageComputeMeter>) {
self.compute_meter = meter;
}
pub(crate) fn active_session_count(&self) -> usize {
self.sessions.len()
}
pub fn kv_pool_tokens(&self) -> u32 {
self.ctx_size
}
#[cfg(test)]
pub(crate) fn is_modelless_for_test(&self) -> bool {
self.modelless_for_test
}
}
impl Drop for RuntimeState {
fn drop(&mut self) {
self.sessions.clear();
self.idle_sessions.clear();
}
}
pub fn load_runtime(config: &StageConfig) -> Result<Option<Arc<Mutex<RuntimeState>>>> {
load_runtime_with_overrides(config, &RuntimeLaunchOverrides::default(), None)
}
pub fn loaded_model_state_kind(
runtime: Option<&Arc<Mutex<RuntimeState>>>,
) -> Option<ModelStateKind> {
runtime.and_then(|runtime| {
runtime
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.model
.capability()
.map(|capability| capability.state_kind)
})
}
pub fn loaded_model_has_indexer_memory(runtime: Option<&Arc<Mutex<RuntimeState>>>) -> Option<bool> {
runtime.and_then(|runtime| {
runtime
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.model
.capability()
.map(|capability| capability.has_indexer_memory)
})
}
pub fn load_runtime_with_overrides(
config: &StageConfig,
overrides: &RuntimeLaunchOverrides,
session_lifecycle_observer: Option<Arc<dyn SessionLifecycleObserver>>,
) -> Result<Option<Arc<Mutex<RuntimeState>>>> {
reject_legacy_serving_package(config)?;
let runtime_config = runtime_config_from_stage_config(config, overrides)?;
let admitted_model_parts = config
.model_part_paths
.iter()
.map(std::path::PathBuf::from)
.collect::<Vec<_>>();
let model = match config.load_mode {
_ if std::env::var("MESH_LLM_BYPASS_SKIPPY_MODEL_LOAD").is_ok() => {
skippy_runtime::StageModel::new_dummy()
}
_ if !admitted_model_parts.is_empty() => {
open_stage_model_from_parts(&admitted_model_parts, &runtime_config)?
}
_ => {
let Some(model_path) = config.model_path.as_ref().map(std::path::Path::new) else {
return Ok(None);
};
open_stage_model(model_path, &runtime_config)?
}
};
Ok(Some(runtime_from_loaded_model(
config,
model,
session_lifecycle_observer,
)?))
}
fn runtime_from_loaded_model(
config: &StageConfig,
model: StageModel,
session_lifecycle_observer: Option<Arc<dyn SessionLifecycleObserver>>,
) -> Result<Arc<Mutex<RuntimeState>>> {
reject_unsupported_staged_workload(config, &model)?;
let lane_count = effective_lane_count(config.lane_count, &model)?;
Ok(Arc::new(Mutex::new(RuntimeState {
model,
layer_start: config.layer_start,
layer_end: config.layer_end,
lane_count,
ctx_size: config.ctx_size,
next_lane_index: 0,
free_lane_indices: Vec::new(),
sessions: BTreeMap::new(),
idle_sessions: Vec::new(),
max_idle_sessions: max_idle_sessions_from_stage_config(config),
session_token_counts: BTreeMap::new(),
session_resident_prefixes: BTreeMap::new(),
session_lifecycle_observer,
compute_meter: Arc::default(),
#[cfg(test)]
modelless_for_test: false,
})))
}
pub fn load_runtime_with_overrides_and_open_events(
config: &StageConfig,
overrides: &RuntimeLaunchOverrides,
model_open_events: Option<&Arc<skippy_runtime::ModelOpenEventQueue>>,
session_lifecycle_observer: Option<Arc<dyn SessionLifecycleObserver>>,
) -> Result<Option<Arc<Mutex<RuntimeState>>>> {
reject_legacy_serving_package(config)?;
let runtime_config = runtime_config_from_stage_config(config, overrides)?;
let admitted_model_parts = config
.model_part_paths
.iter()
.map(std::path::PathBuf::from)
.collect::<Vec<_>>();
let model = match config.load_mode {
_ if std::env::var("MESH_LLM_BYPASS_SKIPPY_MODEL_LOAD").is_ok() => {
skippy_runtime::StageModel::new_dummy()
}
_ if !admitted_model_parts.is_empty() => open_stage_model_from_parts_with_events(
&admitted_model_parts,
&runtime_config,
model_open_events,
)?,
_ => {
let Some(model_path) = config.model_path.as_ref().map(std::path::Path::new) else {
return Ok(None);
};
open_stage_model_with_events(model_path, &runtime_config, model_open_events)?
}
};
Ok(Some(runtime_from_loaded_model(
config,
model,
session_lifecycle_observer,
)?))
}
fn effective_lane_count(configured: u32, model: &StageModel) -> Result<u32> {
if !model.has_native_model() {
return Ok(configured);
}
Ok(lane_count_for_workload(
configured,
model.workload_info()?.kind,
))
}
fn lane_count_for_workload(configured: u32, workload: ModelWorkload) -> u32 {
if configured <= 1 {
return configured;
}
if workload == ModelWorkload::EncoderDecoder {
return 1;
}
configured
}
pub(crate) fn reject_unsupported_staged_workload(
config: &StageConfig,
model: &StageModel,
) -> Result<()> {
if !model.has_native_model()
|| (config.resident_tensor_names.is_empty() && config.layer_start == 0)
{
return Ok(());
}
if model.supports_speech_synthesis() {
anyhow::bail!(
"unsupported staged workload speech_synthesis: audio generation requires an unsplit full model"
);
}
let workload = model.workload_info()?.kind;
if workload != ModelWorkload::CausalGeneration {
anyhow::bail!(
"unsupported staged workload {}: non-chat execution requires an unsplit full model",
match workload {
ModelWorkload::CausalGeneration => unreachable!(),
ModelWorkload::Embedding => "embedding",
ModelWorkload::Rerank => "rerank",
ModelWorkload::EncoderDecoder => "encoder_decoder",
}
);
}
Ok(())
}
fn max_idle_sessions_from_stage_config(config: &StageConfig) -> Option<usize> {
config.cache_idle_slots.map(|slots| slots as usize)
}
fn reject_legacy_serving_package(config: &StageConfig) -> Result<()> {
anyhow::ensure!(
config.load_mode != LoadMode::LayerPackage,
"layer-package schema v1 is offline-only; split serving requires package-v2 graph admission"
);
Ok(())
}
fn runtime_config_from_stage_config(
config: &StageConfig,
overrides: &RuntimeLaunchOverrides,
) -> Result<RuntimeConfig> {
let cache_type_k = parse_cache_type(&config.cache_type_k)
.with_context(|| format!("parse cache_type_k for {}", config.stage_id))?;
let cache_type_v = parse_cache_type(&config.cache_type_v)
.with_context(|| format!("parse cache_type_v for {}", config.stage_id))?;
let n_threads = overrides
.n_threads
.map(u32::try_from)
.transpose()
.with_context(|| format!("n_threads exceeds u32 for {}", config.stage_id))?;
let n_threads_batch = overrides
.n_threads_batch
.map(u32::try_from)
.transpose()
.with_context(|| format!("n_threads_batch exceeds u32 for {}", config.stage_id))?;
Ok(RuntimeConfig {
stage_index: config.stage_index,
layer_start: config.layer_start,
layer_end: config.layer_end,
ctx_size: config.ctx_size,
lane_count: config.lane_count,
n_batch: config.n_batch,
n_ubatch: config.n_ubatch,
n_threads,
n_threads_batch,
n_gpu_layers: config.n_gpu_layers,
mmap: config.mmap,
mlock: config.mlock,
repack: config.repack,
op_offload: config.op_offload,
no_host_buffer: config.no_host_buffer,
check_tensors: config.check_tensors,
direct_io: config.direct_io,
main_gpu: config.main_gpu,
split_mode: match config.split_mode {
SplitMode::Auto => RuntimeSplitMode::Auto,
SplitMode::None => RuntimeSplitMode::None,
SplitMode::Layer => RuntimeSplitMode::Layer,
SplitMode::Row => RuntimeSplitMode::Row,
SplitMode::Tensor => RuntimeSplitMode::Tensor,
},
selected_backend_device: config
.selected_device
.as_ref()
.map(|device| device.backend_device.clone()),
cache_type_k,
cache_type_v,
flash_attn_type: match config.flash_attn_type {
FlashAttentionType::Auto => RuntimeFlashAttentionType::Auto,
FlashAttentionType::Disabled => RuntimeFlashAttentionType::Disabled,
FlashAttentionType::Enabled => RuntimeFlashAttentionType::Enabled,
},
load_mode: match config.load_mode {
LoadMode::RuntimeSlice => RuntimeLoadMode::RuntimeSlice,
LoadMode::LayerPackage => RuntimeLoadMode::LayerPackage,
LoadMode::ArtifactSlice => RuntimeLoadMode::ArtifactSlice,
},
kv_offload: config.kv_offload,
kv_unified: config.kv_unified,
swa_full: config.swa_full,
projector_path: config.projector_path.clone(),
projector_use_gpu: config.projector_use_gpu,
media_marker: config.media_marker.clone(),
image_min_tokens: config.image_min_tokens,
image_max_tokens: config.image_max_tokens,
batch_max_tokens: config.batch_max_tokens,
glm_dsa_policy: match config.glm_dsa_policy {
skippy_protocol::GlmDsaPolicy::Auto => RuntimeGlmDsaPolicy::Auto,
skippy_protocol::GlmDsaPolicy::V1 => RuntimeGlmDsaPolicy::V1,
},
mtp_source: overrides.mtp_source,
resident_tensor_names: config.resident_tensor_names.clone(),
execution_contract: config.execution_contract.clone(),
activation_import_identities: config.activation_import_identities.clone(),
activation_import_bindings: config.activation_import_bindings.clone(),
activation_export_identities: config.activation_export_identities.clone(),
activation_export_bindings: config.activation_export_bindings.clone(),
checkpoint_quantization: config
.checkpoint_quantization
.as_deref()
.unwrap_or("preserve")
.parse()
.map_err(anyhow::Error::msg)
.with_context(|| format!("parse checkpoint_quantization for {}", config.stage_id))?,
checkpoint_imatrix: config.checkpoint_imatrix.as_deref().map(Into::into),
checkpoint_imatrix_sha256: config.checkpoint_imatrix_sha256.clone(),
})
}
fn open_stage_model(path: &std::path::Path, runtime_config: &RuntimeConfig) -> Result<StageModel> {
StageModel::open(path, runtime_config)
}
fn open_stage_model_with_events(
path: &std::path::Path,
runtime_config: &RuntimeConfig,
model_open_events: Option<&Arc<skippy_runtime::ModelOpenEventQueue>>,
) -> Result<StageModel> {
match model_open_events {
Some(queue) => StageModel::open_with_events(path, runtime_config, queue),
None => StageModel::open(path, runtime_config),
}
}
fn open_stage_model_from_parts(
paths: &[std::path::PathBuf],
runtime_config: &RuntimeConfig,
) -> Result<StageModel> {
StageModel::open_from_parts(paths, runtime_config)
}
fn open_stage_model_from_parts_with_events(
paths: &[std::path::PathBuf],
runtime_config: &RuntimeConfig,
model_open_events: Option<&Arc<skippy_runtime::ModelOpenEventQueue>>,
) -> Result<StageModel> {
match model_open_events {
Some(queue) => StageModel::open_from_parts_with_events(paths, runtime_config, queue),
None => StageModel::open_from_parts(paths, runtime_config),
}
}
#[cfg(test)]
mod tests {
use skippy_protocol::{
FlashAttentionType, LoadMode, PeerConfig, SplitMode, StageConfig, StageDevice,
};
use skippy_runtime::{
ActivationFrame, CheckpointQuantization, FlashAttentionType as RuntimeFlashAttentionType,
ModelWorkload, MtpSource, RuntimeConfig, SamplingConfig, StageModel,
};
use super::{
RuntimeLaunchOverrides, RuntimeState, effective_lane_count, lane_count_for_workload,
load_runtime_with_overrides, max_idle_sessions_from_stage_config,
reject_legacy_serving_package, runtime_config_from_stage_config, runtime_from_loaded_model,
};
#[test]
fn filtered_dummy_model_retains_runtime_load_bypass() {
let config = StageConfig {
resident_tensor_names: vec!["blk.0.attn_q.weight".to_owned()],
lane_count: 2,
..Default::default()
};
let runtime = super::runtime_from_loaded_model(
&config,
skippy_runtime::StageModel::new_dummy(),
None,
)
.unwrap();
let runtime = runtime.lock().unwrap();
assert!(!runtime.model.has_native_model());
assert_eq!(runtime.lane_count(), 2);
}
#[test]
fn workload_admission_constructor_initializes_shared_compute_meter() {
let runtime =
runtime_from_loaded_model(&StageConfig::default(), StageModel::new_dummy(), None)
.expect("dummy construction succeeds");
let runtime = runtime.lock().unwrap();
let meter = runtime.compute_meter();
assert_eq!(
meter.snapshot(),
crate::compute_meter::StageComputeSnapshot::default()
);
meter.record(std::time::Duration::from_millis(2));
meter.record_decode_tokens(3);
assert_eq!(
runtime.compute_meter().snapshot(),
crate::compute_meter::StageComputeSnapshot {
busy_nanos: 2_000_000,
operations: 1,
decode_tokens: 3,
}
);
}
#[test]
fn modelless_runtime_reports_zero_kv_pool_so_scheduler_uses_fallback() {
let rt = RuntimeState::new_modelless_for_test(4);
assert_eq!(rt.kv_pool_tokens(), 0);
assert_eq!(rt.lane_count(), 4);
}
#[test]
fn encoder_decoder_lane_admission_clamps_to_the_native_single_lane() {
assert_eq!(lane_count_for_workload(4, ModelWorkload::EncoderDecoder), 1);
assert_eq!(lane_count_for_workload(1, ModelWorkload::EncoderDecoder), 1);
assert_eq!(
lane_count_for_workload(4, ModelWorkload::CausalGeneration),
4
);
assert_eq!(lane_count_for_workload(4, ModelWorkload::Embedding), 4);
assert_eq!(lane_count_for_workload(4, ModelWorkload::Rerank), 4);
}
#[test]
fn load_bypass_dummy_models_keep_the_configured_lane_count() {
assert_eq!(
effective_lane_count(4, &StageModel::new_dummy()).unwrap(),
4
);
let config = StageConfig {
stage_id: "stage-0".to_string(),
lane_count: 4,
..StageConfig::default()
};
let runtime = runtime_from_loaded_model(&config, StageModel::new_dummy(), None)
.expect("dummy construction succeeds");
let rt = runtime
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
assert_eq!(rt.lane_count(), 4);
}
#[test]
fn runtime_config_preserves_selected_backend_device_and_thread_overrides() {
let config = StageConfig {
run_id: "run-a".to_string(),
topology_id: "topology-a".to_string(),
model_id: "model-a".to_string(),
package_ref: None,
manifest_sha256: None,
source_model_path: None,
source_model_sha256: None,
source_model_bytes: None,
materialized_path: None,
materialized_pinned: false,
model_path: Some("/tmp/model.gguf".to_string()),
projector_path: Some("/tmp/mmproj.gguf".to_string()),
stage_id: "stage-0".to_string(),
stage_index: 0,
layer_start: 0,
layer_end: 24,
ctx_size: 512,
lane_count: 2,
n_batch: Some(1024),
n_ubatch: Some(256),
n_gpu_layers: -1,
mmap: Some(false),
mlock: true,
repack: false,
op_offload: None,
no_host_buffer: false,
check_tensors: false,
direct_io: false,
main_gpu: None,
split_mode: SplitMode::Auto,
cache_type_k: "f16".to_string(),
cache_type_v: "f16".to_string(),
flash_attn_type: FlashAttentionType::Enabled,
kv_offload: None,
kv_unified: None,
swa_full: None,
cache_idle_slots: None,
resident_tensor_names: Vec::new(),
selected_device: Some(StageDevice {
backend_device: "Vulkan1".into(),
stable_id: Some("pci:0000:65:00.0".into()),
index: Some(1),
vram_bytes: Some(16_000_000_000),
}),
kv_cache: None,
native_mtp_enabled: true,
load_mode: LoadMode::RuntimeSlice,
bind_addr: "127.0.0.1:0".to_string(),
upstream: None,
downstream: None,
..StageConfig::default()
};
let overrides = RuntimeLaunchOverrides {
n_threads: Some(8),
n_threads_batch: Some(4),
mtp_source: MtpSource::External,
};
let runtime_config = runtime_config_from_stage_config(&config, &overrides).unwrap();
assert_eq!(
runtime_config.selected_backend_device.as_deref(),
Some("Vulkan1")
);
assert_eq!(runtime_config.lane_count, 2);
assert_eq!(runtime_config.n_batch, Some(1024));
assert_eq!(runtime_config.n_ubatch, Some(256));
assert_eq!(runtime_config.n_threads, Some(8));
assert_eq!(runtime_config.n_threads_batch, Some(4));
assert_eq!(runtime_config.mmap, Some(false));
assert!(runtime_config.mlock);
assert_eq!(
runtime_config.flash_attn_type,
RuntimeFlashAttentionType::Enabled
);
assert_eq!(runtime_config.mtp_source, MtpSource::External);
}
#[test]
fn runtime_config_parses_checkpoint_quantization() {
let config = StageConfig {
stage_id: "stage-0".to_string(),
layer_end: 1,
checkpoint_quantization: Some("Q4_K_M".to_string()),
..StageConfig::default()
};
let runtime_config =
runtime_config_from_stage_config(&config, &RuntimeLaunchOverrides::default()).unwrap();
assert_eq!(
runtime_config.checkpoint_quantization,
CheckpointQuantization::Q4KM
);
}
#[test]
fn runtime_config_accepts_importance_aware_checkpoint_quantization() {
let config = StageConfig {
stage_id: "stage-0".to_string(),
layer_end: 1,
checkpoint_quantization: Some("IQ2_XXS".to_string()),
checkpoint_imatrix: Some("/models/imatrix.gguf".to_string()),
checkpoint_imatrix_sha256: Some("a".repeat(64)),
..StageConfig::default()
};
let runtime =
runtime_config_from_stage_config(&config, &RuntimeLaunchOverrides::default()).unwrap();
assert_eq!(
runtime.checkpoint_quantization,
CheckpointQuantization::IQ2XXS
);
assert_eq!(
runtime.checkpoint_imatrix.as_deref(),
Some(std::path::Path::new("/models/imatrix.gguf"))
);
assert_eq!(runtime.checkpoint_imatrix_sha256, Some("a".repeat(64)));
}
fn fake_stage_config_with_cache_idle_slots(cache_idle_slots: Option<u32>) -> StageConfig {
StageConfig {
run_id: "run-a".to_string(),
topology_id: "topology-a".to_string(),
model_id: "model-a".to_string(),
package_ref: None,
manifest_sha256: None,
source_model_path: None,
source_model_sha256: None,
source_model_bytes: None,
materialized_path: None,
materialized_pinned: false,
model_path: Some("/tmp/model.gguf".to_string()),
projector_path: None,
stage_id: "stage-0".to_string(),
stage_index: 0,
layer_start: 0,
layer_end: 24,
ctx_size: 512,
lane_count: 4,
n_batch: None,
n_ubatch: None,
n_gpu_layers: -1,
mmap: None,
mlock: false,
repack: false,
op_offload: None,
no_host_buffer: false,
check_tensors: false,
direct_io: false,
main_gpu: None,
split_mode: SplitMode::Auto,
cache_type_k: "f16".to_string(),
cache_type_v: "f16".to_string(),
flash_attn_type: FlashAttentionType::Auto,
kv_offload: None,
kv_unified: None,
swa_full: None,
cache_idle_slots,
resident_tensor_names: Vec::new(),
selected_device: None,
kv_cache: None,
native_mtp_enabled: true,
load_mode: LoadMode::RuntimeSlice,
bind_addr: "127.0.0.1:0".to_string(),
upstream: None,
downstream: None,
..StageConfig::default()
}
}
#[test]
fn cache_idle_slots_reaches_the_idle_session_pool_bound() {
let unset = fake_stage_config_with_cache_idle_slots(None);
let two = fake_stage_config_with_cache_idle_slots(Some(2));
let five = fake_stage_config_with_cache_idle_slots(Some(5));
assert_eq!(max_idle_sessions_from_stage_config(&unset), None);
assert_eq!(max_idle_sessions_from_stage_config(&two), Some(2));
assert_eq!(max_idle_sessions_from_stage_config(&five), Some(5));
assert_ne!(
max_idle_sessions_from_stage_config(&two),
max_idle_sessions_from_stage_config(&five),
"cache_idle_slots=2 and cache_idle_slots=5 must produce different idle-pool bounds"
);
}
#[test]
fn legacy_layer_packages_are_rejected_before_model_open() {
let mut config = fake_stage_config_with_cache_idle_slots(None);
config.load_mode = LoadMode::LayerPackage;
let error = reject_legacy_serving_package(&config)
.expect_err("schema-v1 layer packages must be offline-only");
assert!(error.to_string().contains("package-v2 graph admission"));
}
#[test]
fn runtime_config_does_not_infer_embedding_ownership_from_stage_role() {
let config = StageConfig {
run_id: "run-a".to_string(),
topology_id: "topology-a".to_string(),
model_id: "model-a".to_string(),
package_ref: Some("/tmp/package".to_string()),
manifest_sha256: Some("manifest".to_string()),
source_model_path: None,
source_model_sha256: None,
source_model_bytes: None,
materialized_path: None,
materialized_pinned: false,
model_path: Some("/tmp/package".to_string()),
projector_path: None,
stage_id: "stage-2".to_string(),
stage_index: 2,
layer_start: 20,
layer_end: 30,
ctx_size: 512,
lane_count: 1,
n_batch: None,
n_ubatch: None,
n_gpu_layers: -1,
mmap: None,
mlock: false,
repack: false,
op_offload: None,
no_host_buffer: false,
check_tensors: false,
direct_io: false,
main_gpu: None,
split_mode: SplitMode::Auto,
cache_type_k: "f16".to_string(),
cache_type_v: "f16".to_string(),
flash_attn_type: FlashAttentionType::Auto,
kv_offload: None,
kv_unified: None,
swa_full: None,
cache_idle_slots: None,
resident_tensor_names: Vec::new(),
selected_device: Some(StageDevice {
backend_device: "CPU".into(),
stable_id: None,
index: None,
vram_bytes: None,
}),
kv_cache: None,
native_mtp_enabled: true,
load_mode: LoadMode::LayerPackage,
bind_addr: "127.0.0.1:0".to_string(),
upstream: Some(PeerConfig {
stage_id: "stage-1".to_string(),
stage_index: 1,
endpoint: "tcp://127.0.0.1:19001".to_string(),
}),
downstream: None,
..StageConfig::default()
};
let runtime_config =
runtime_config_from_stage_config(&config, &RuntimeLaunchOverrides::default()).unwrap();
assert!(runtime_config.is_terminal_stage());
assert_eq!(runtime_config.mtp_source, MtpSource::Disabled);
}
fn glm52_mtp_fixture() -> anyhow::Result<Option<(std::path::PathBuf, StageConfig)>> {
let Some(package_path) =
std::env::var_os("SKIPPY_GLM52_MTP_PACKAGE").map(std::path::PathBuf::from)
else {
return Ok(None);
};
if !package_path.join("model-package.json").is_file() {
eprintln!(
"skipping: {} does not look like a package-v2 directory",
package_path.display()
);
return Ok(None);
}
let manifest_path = package_path.join("model-package.json");
let manifest: skippy_package_format::PackageManifest =
serde_json::from_slice(&std::fs::read(&manifest_path)?)?;
let descriptor = skippy_package_format::stage_admission::StageAdmissionDescriptor {
package_id: manifest.package_id.clone(),
resident_tensor_ids: manifest
.tensor_catalog
.entries
.iter()
.filter(|tensor| {
matches!(
tensor.storage,
skippy_package_format::TensorStorage::Owned { .. }
)
})
.map(|tensor| tensor.id.clone())
.collect(),
sidecars: Vec::new(),
};
let admission = manifest.resolve_stage_admission(&descriptor)?;
let mut resident_tensor_names = admission
.tensor_bindings
.iter()
.map(|tensor| tensor.native_name.to_string())
.collect::<Vec<_>>();
resident_tensor_names.sort();
resident_tensor_names.dedup();
let mut artifacts = admission.required_artifacts;
artifacts.sort_by(|left, right| {
let left_primary = left.id == manifest.source_model.metadata_artifact_id;
let right_primary = right.id == manifest.source_model.metadata_artifact_id;
right_primary
.cmp(&left_primary)
.then_with(|| left.id.cmp(&right.id))
});
let model_part_paths = artifacts
.into_iter()
.map(|artifact| {
package_path
.join(&artifact.path)
.to_string_lossy()
.into_owned()
})
.collect::<Vec<_>>();
let model_path = model_part_paths.first().cloned();
let config = StageConfig {
run_id: "glm52-mtp-smoke".to_string(),
topology_id: "glm52-mtp-smoke-topology".to_string(),
model_id: "meshllm/GLM-5.2-Q2_K-MTP-Q8-layers".to_string(),
package_ref: Some(package_path.to_string_lossy().to_string()),
manifest_sha256: None,
source_model_path: None,
source_model_sha256: None,
source_model_bytes: None,
materialized_path: None,
materialized_pinned: false,
model_path,
model_part_paths,
projector_path: None,
stage_id: "stage-final".to_string(),
stage_index: 1,
layer_start: 74,
layer_end: 78,
ctx_size: 128,
lane_count: 1,
n_batch: Some(1),
n_ubatch: Some(1),
n_gpu_layers: 0,
mmap: Some(true),
mlock: false,
repack: false,
op_offload: None,
no_host_buffer: false,
check_tensors: false,
direct_io: false,
main_gpu: None,
split_mode: SplitMode::Auto,
cache_type_k: "f16".to_string(),
cache_type_v: "f16".to_string(),
flash_attn_type: FlashAttentionType::Disabled,
kv_offload: None,
kv_unified: None,
swa_full: None,
cache_idle_slots: None,
resident_tensor_names,
selected_device: Some(StageDevice {
backend_device: "CPU".into(),
stable_id: None,
index: None,
vram_bytes: None,
}),
kv_cache: None,
native_mtp_enabled: true,
load_mode: LoadMode::RuntimeSlice,
bind_addr: "127.0.0.1:0".to_string(),
upstream: Some(PeerConfig {
stage_id: "stage-prev".to_string(),
stage_index: 0,
endpoint: "tcp://127.0.0.1:19000".to_string(),
}),
downstream: None,
..StageConfig::default()
};
Ok(Some((package_path, config)))
}
fn glm52_mtp_input(token_count: u32) -> ActivationFrame {
let hidden_bytes = 6144 * token_count as usize * std::mem::size_of::<f32>();
let mut frame = crate::test_activation::frame(
token_count,
vec![crate::test_activation::PartBytes {
identity: 1,
ggml_type: skippy_runtime::GGML_TYPE_F32,
flags: 0,
bytes: vec![0; hidden_bytes],
}],
);
frame.desc.producer_stage_index = 0;
frame.desc.layer_start = 0;
frame.desc.layer_end = 74;
frame
}
#[test]
fn glm52_final_stage_package_executes_native_mtp_when_fixture_is_set() -> anyhow::Result<()> {
let Some((_package_path, config)) = glm52_mtp_fixture()? else {
eprintln!("skipping: SKIPPY_GLM52_MTP_PACKAGE is not set");
return Ok(());
};
let runtime = load_runtime_with_overrides(
&config,
&RuntimeLaunchOverrides {
mtp_source: MtpSource::Integrated,
..RuntimeLaunchOverrides::default()
},
None,
)?
.expect("GLM final stage should load from the package");
let mut runtime = runtime.lock().expect("runtime mutex poisoned");
let input = glm52_mtp_input(1);
let sampling = SamplingConfig {
temperature: 0.0,
..SamplingConfig::default()
};
let (predicted, draft, _output) =
runtime.decode_frame_sampled_mtp("smoke", 1, Some(&sampling), Some(&input), 0, 1)?;
let draft = draft.expect("GLM final stage should return a native MTP draft");
assert!(predicted >= 0);
assert_eq!(draft.token_ids.len(), 1);
assert!(draft.token_ids[0] >= 0);
let verify_inputs = [predicted, draft.token_ids[0]];
let (verified, _next_draft, _output) = runtime.verify_frame_sampled(
"smoke",
&verify_inputs,
Some(&sampling),
Some(&glm52_mtp_input(2)),
0,
1,
)?;
assert!(!verified.is_empty());
runtime.retire_verify_checkpoint("smoke", 1, 2)?;
Ok(())
}
#[test]
fn glm52_final_stage_does_not_create_integrated_mtp_when_disabled() -> anyhow::Result<()> {
let Some((_package_path, config)) = glm52_mtp_fixture()? else {
eprintln!("skipping: SKIPPY_GLM52_MTP_PACKAGE is not set");
return Ok(());
};
let runtime = load_runtime_with_overrides(
&config,
&RuntimeLaunchOverrides {
mtp_source: MtpSource::Disabled,
..RuntimeLaunchOverrides::default()
},
None,
)?
.expect("GLM final stage should load from the package");
let mut runtime = runtime.lock().expect("runtime mutex poisoned");
let sampling = SamplingConfig {
temperature: 0.0,
..SamplingConfig::default()
};
let (predicted, draft, _output) = runtime.decode_frame_sampled_mtp(
"disabled-mtp",
1,
Some(&sampling),
Some(&glm52_mtp_input(1)),
0,
1,
)?;
assert!(predicted >= 0);
assert!(
draft.is_none(),
"disabled MTP must not create a draft context"
);
Ok(())
}
#[test]
fn glm52_external_sidecar_attaches_when_target_has_integrated_mtp_tensors() -> anyhow::Result<()>
{
let Some((_package_path, config)) = glm52_mtp_fixture()? else {
eprintln!("skipping: SKIPPY_GLM52_MTP_PACKAGE is not set");
return Ok(());
};
let Some(sidecar_path) = std::env::var_os("SKIPPY_GLM52_MTP_SIDECAR") else {
eprintln!("skipping: SKIPPY_GLM52_MTP_SIDECAR is not set");
return Ok(());
};
let sidecar_path = std::path::PathBuf::from(sidecar_path);
if !sidecar_path.is_file() {
eprintln!(
"skipping: {} is not an MTP sidecar GGUF",
sidecar_path.display()
);
return Ok(());
}
let runtime = load_runtime_with_overrides(
&config,
&RuntimeLaunchOverrides {
mtp_source: MtpSource::External,
..RuntimeLaunchOverrides::default()
},
None,
)?
.expect("GLM final stage should load from the package");
let mut runtime = runtime.lock().expect("runtime mutex poisoned");
runtime.model.attach_mtp_draft_model(
&sidecar_path,
&RuntimeConfig {
ctx_size: config.ctx_size,
lane_count: config.lane_count,
n_batch: config.n_batch,
n_ubatch: config.n_ubatch,
n_gpu_layers: config.n_gpu_layers,
mmap: config.mmap,
mlock: config.mlock,
selected_backend_device: Some("CPU".to_string()),
mtp_source: MtpSource::External,
..RuntimeConfig::default()
},
)?;
let sampling = SamplingConfig {
temperature: 0.0,
..SamplingConfig::default()
};
let (predicted, draft, _output) = runtime.decode_frame_sampled_mtp(
"external-mtp",
1,
Some(&sampling),
Some(&glm52_mtp_input(1)),
0,
1,
)?;
let draft = draft.expect("external MTP sidecar should attach to the target");
assert!(predicted >= 0);
assert_eq!(draft.token_ids.len(), 1);
assert!(draft.token_ids[0] >= 0);
Ok(())
}
#[test]
fn runtime_config_preserves_default_runtime_threads_when_omitted() {
let config = StageConfig {
run_id: "run-a".to_string(),
topology_id: "topology-a".to_string(),
model_id: "model-a".to_string(),
package_ref: None,
manifest_sha256: None,
source_model_path: None,
source_model_sha256: None,
source_model_bytes: None,
materialized_path: None,
materialized_pinned: false,
model_path: Some("/tmp/model.gguf".to_string()),
projector_path: None,
stage_id: "stage-0".to_string(),
stage_index: 0,
layer_start: 0,
layer_end: 24,
ctx_size: 512,
lane_count: 1,
n_batch: None,
n_ubatch: None,
n_gpu_layers: -1,
mmap: None,
mlock: false,
repack: false,
op_offload: None,
no_host_buffer: false,
check_tensors: false,
direct_io: false,
main_gpu: None,
split_mode: SplitMode::Auto,
cache_type_k: "f16".to_string(),
cache_type_v: "f16".to_string(),
flash_attn_type: FlashAttentionType::Auto,
kv_offload: None,
kv_unified: None,
swa_full: None,
cache_idle_slots: None,
resident_tensor_names: Vec::new(),
selected_device: None,
kv_cache: None,
native_mtp_enabled: true,
load_mode: LoadMode::RuntimeSlice,
bind_addr: "127.0.0.1:0".to_string(),
upstream: None,
downstream: None,
..StageConfig::default()
};
let runtime_config =
runtime_config_from_stage_config(&config, &RuntimeLaunchOverrides::default()).unwrap();
assert_eq!(runtime_config.n_threads, None);
assert_eq!(runtime_config.n_threads_batch, None);
assert_eq!(runtime_config.n_batch, None);
assert_eq!(runtime_config.n_ubatch, None);
}
#[test]
fn runtime_config_preserves_multimodal_and_glm_dsa_native_controls() {
let config: StageConfig = serde_json::from_value(serde_json::json!({
"run_id": "run-a",
"topology_id": "topology-a",
"model_id": "model-a",
"model_path": "/tmp/model.gguf",
"projector_path": "/tmp/mmproj.gguf",
"projector_use_gpu": false,
"media_marker": "<media>",
"image_min_tokens": 32,
"image_max_tokens": 1536,
"batch_max_tokens": 384,
"glm_dsa_policy": "v1",
"generation_signal_window": 20,
"stage_id": "stage-0",
"stage_index": 0,
"layer_start": 0,
"layer_end": 24,
"ctx_size": 512,
"lane_count": 1,
"n_gpu_layers": -1,
"cache_type_k": "f16",
"cache_type_v": "f16",
"native_mtp_enabled": true,
"load_mode": "runtime-slice",
"execution_contract": "",
"bind_addr": "127.0.0.1:0"
}))
.expect("stage config should deserialize");
let runtime_config =
runtime_config_from_stage_config(&config, &RuntimeLaunchOverrides::default())
.expect("runtime config should build");
let debug = format!("{runtime_config:?}");
assert!(debug.contains("projector_use_gpu: Some(false)"));
assert!(debug.contains("media_marker: Some(\"<media>\")"));
assert!(debug.contains("image_min_tokens: Some(32)"));
assert!(debug.contains("image_max_tokens: Some(1536)"));
assert!(debug.contains("batch_max_tokens: Some(384)"));
assert!(debug.contains("glm_dsa_policy: V1"));
}
#[test]
fn runtime_config_rejects_unsupported_cache_type_before_launch() {
let config = StageConfig {
run_id: "run-a".to_string(),
topology_id: "topology-a".to_string(),
model_id: "model-a".to_string(),
package_ref: None,
manifest_sha256: None,
source_model_path: None,
source_model_sha256: None,
source_model_bytes: None,
materialized_path: None,
materialized_pinned: false,
model_path: Some("/tmp/model.gguf".to_string()),
projector_path: None,
stage_id: "stage-0".to_string(),
stage_index: 0,
layer_start: 0,
layer_end: 24,
ctx_size: 512,
lane_count: 1,
n_batch: None,
n_ubatch: None,
n_gpu_layers: -1,
mmap: None,
mlock: false,
repack: false,
op_offload: None,
no_host_buffer: false,
check_tensors: false,
direct_io: false,
main_gpu: None,
split_mode: SplitMode::Auto,
cache_type_k: "auto".to_string(),
cache_type_v: "f16".to_string(),
flash_attn_type: FlashAttentionType::Auto,
kv_offload: None,
kv_unified: None,
swa_full: None,
cache_idle_slots: None,
resident_tensor_names: Vec::new(),
selected_device: None,
kv_cache: None,
native_mtp_enabled: true,
load_mode: LoadMode::RuntimeSlice,
bind_addr: "127.0.0.1:0".to_string(),
upstream: None,
downstream: None,
..StageConfig::default()
};
let error = runtime_config_from_stage_config(&config, &RuntimeLaunchOverrides::default())
.expect_err("unsupported cache types should fail during runtime config construction");
assert!(
error.to_string().contains("parse cache_type_k for stage-0"),
"unexpected error: {error:#}"
);
}
}