use std::sync::Arc;
use std::time::Instant;
use anyhow::{Context, Result};
use skippy_cache::ExactStatePayload;
use crate::runtime_state::RuntimeState;
use super::{
ExactStateExtra, ExactStateRecord, ExactStateRecordAdmission, ExactStateRestore,
KvStageIntegration, PendingExactStateRecord, PrefillKvIdentity, StagePrefixCachePayload,
records::add_reconstruct_stats,
};
const MAX_PREFIX_PROBES: usize = 64;
fn l3_fill_claim_key(l3: &skippy_cache::L3Tier, location: &skippy_cache::L3Location) -> String {
format!("{}:{}", l3.state_identity(), location.manifest_key)
}
fn resident_prefix_is_preferred(matched_tokens: usize, exact_tokens: Option<usize>) -> bool {
exact_tokens.is_none_or(|exact_tokens| matched_tokens > exact_tokens)
}
fn preflight_l3_kv_location(
location: &skippy_cache::L3Location,
) -> Result<Option<skippy_runtime::RuntimeKvPageDesc>> {
if !location.native_kv_passthrough && !location.cachegen_kv {
return Ok(None);
}
let json = location
.kv_desc_json
.as_deref()
.context("native KV manifest has no runtime page descriptor")?;
let desc: skippy_runtime::RuntimeKvPageDesc =
serde_json::from_str(json).context("native KV manifest has an invalid page descriptor")?;
let declared_bytes = if location.cachegen_kv {
location.kv_decoded_bytes
} else {
location.kv_bytes
};
let kv_bytes = usize::try_from(declared_bytes).context("KV payload length exceeds usize")?;
desc.validate_payload(kv_bytes)
.context("native KV manifest page descriptor is incompatible")?;
if desc.token_start != 0 || desc.token_count != location.token_count {
anyhow::bail!("native KV manifest page descriptor does not cover the located prefix");
}
Ok(Some(desc))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CaptureAdmission {
Continuation,
BestEffort,
}
pub const EXACT_STATE_ADMISSION_TOTAL: usize = 2;
pub const EXACT_STATE_ADMISSION_RESERVED: usize = 1;
#[derive(Debug)]
pub struct ExactStateAdmissionCredit {
class: CaptureAdmission,
total: Arc<std::sync::atomic::AtomicUsize>,
best_effort: Arc<std::sync::atomic::AtomicUsize>,
}
impl ExactStateAdmissionCredit {
pub(crate) fn acquire(
total: &Arc<std::sync::atomic::AtomicUsize>,
best_effort: &Arc<std::sync::atomic::AtomicUsize>,
class: CaptureAdmission,
) -> Option<Self> {
use std::sync::atomic::Ordering::{AcqRel, Acquire};
match class {
CaptureAdmission::Continuation => {
let mut observed = total.load(Acquire);
loop {
if observed >= EXACT_STATE_ADMISSION_TOTAL {
return None;
}
match total.compare_exchange_weak(observed, observed + 1, AcqRel, Acquire) {
Ok(_) => {
return Some(Self {
class,
total: Arc::clone(total),
best_effort: Arc::clone(best_effort),
});
}
Err(actual) => observed = actual,
}
}
}
CaptureAdmission::BestEffort => {
let reserved_cap = EXACT_STATE_ADMISSION_TOTAL - EXACT_STATE_ADMISSION_RESERVED;
let mut observed_best_effort = best_effort.load(Acquire);
loop {
if observed_best_effort >= reserved_cap {
return None;
}
match best_effort.compare_exchange_weak(
observed_best_effort,
observed_best_effort + 1,
AcqRel,
Acquire,
) {
Ok(_) => break,
Err(actual) => observed_best_effort = actual,
}
}
let mut observed_total = total.load(Acquire);
loop {
if observed_total >= EXACT_STATE_ADMISSION_TOTAL {
best_effort.fetch_sub(1, AcqRel);
return None;
}
match total.compare_exchange_weak(
observed_total,
observed_total + 1,
AcqRel,
Acquire,
) {
Ok(_) => {
return Some(Self {
class,
total: Arc::clone(total),
best_effort: Arc::clone(best_effort),
});
}
Err(actual) => observed_total = actual,
}
}
}
}
}
}
impl Drop for ExactStateAdmissionCredit {
fn drop(&mut self) {
let _ = self.total.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
if self.class == CaptureAdmission::BestEffort {
let _ = self
.best_effort
.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
}
}
}
impl KvStageIntegration {
fn available_exact_prefix_tokens(
&self,
identities: &[PrefillKvIdentity],
shared_durable_fallback: bool,
) -> Option<usize> {
identities
.iter()
.filter_map(|identity| {
let warm = self
.radix
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.peek_recurrent(&identity.namespace, &identity.token_ids)
.map(|exact| exact.matched_tokens);
if warm.is_some() {
return warm;
}
let l3 = self.l3.as_ref()?;
let durable_token_ids = if shared_durable_fallback {
self.durable_exact_lookup_token_ids(&identity.token_ids)
} else {
&identity.token_ids
};
l3.locate_longest(&identity.namespace, durable_token_ids, MAX_PREFIX_PROBES)
.ok()
.flatten()
.and_then(|location| usize::try_from(location.token_count).ok())
})
.max()
}
pub(crate) fn l3_benefit_cost(
&self,
cold_prefill_cost: Option<f64>,
) -> Option<skippy_cache::policy::CostSample> {
self.l3
.as_ref()
.and_then(|l3| l3.benefit_candidate_cost(cold_prefill_cost))
}
pub fn restore_exact_state(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identities: &[PrefillKvIdentity],
) -> Result<Option<ExactStateRestore>> {
self.restore_exact_state_with_cold_cost(runtime, session_id, identities, None)
}
pub(crate) fn restore_exact_state_with_shared_durable_fallback(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identities: &[PrefillKvIdentity],
) -> Result<Option<ExactStateRestore>> {
runtime.restore_transaction(session_id, |runtime| {
self.restore_exact_state_inner(runtime, session_id, identities, None, true)
})
}
pub fn restore_exact_state_with_cold_cost(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identities: &[PrefillKvIdentity],
cold_prefill_cost: Option<f64>,
) -> Result<Option<ExactStateRestore>> {
runtime.restore_transaction(session_id, |runtime| {
self.restore_exact_state_inner(
runtime,
session_id,
identities,
cold_prefill_cost,
false,
)
})
}
fn restore_exact_state_inner(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identities: &[PrefillKvIdentity],
cold_prefill_cost: Option<f64>,
shared_durable_fallback: bool,
) -> Result<Option<ExactStateRestore>> {
if !self.should_lookup() || self.exact_state_payload().is_none() {
return Ok(None);
}
if self.payload == StagePrefixCachePayload::ResidentKv {
let resident_tokens = identities
.iter()
.filter_map(|identity| self.probe_resident_prefix(identity))
.map(|resident| resident.token_count)
.max();
if let Some(resident_tokens) = resident_tokens {
let exact_tokens =
self.available_exact_prefix_tokens(identities, shared_durable_fallback);
if resident_prefix_is_preferred(resident_tokens, exact_tokens) {
return Ok(None);
}
}
}
for identity in identities {
let lookup_started = Instant::now();
let (lookup, entries) = {
let mut radix = self
.radix
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let lookup = radix.acquire_recurrent(&identity.namespace, &identity.token_ids);
let entries = radix.stats().recurrent_entries;
(lookup, entries)
};
#[cfg(test)]
crate::frontend::capture_trace::log_lookup_candidate(
&identity.namespace,
&identity.token_ids,
lookup.is_some(),
lookup.as_ref().map(|matched| matched.stored_tokens.len()),
);
let Some(lookup) = lookup else {
if let Some(restored) = self.restore_from_l3(
runtime,
session_id,
identity,
lookup_started,
cold_prefill_cost,
shared_durable_fallback,
)? {
return Ok(Some(restored));
}
continue;
};
let lease = ExactStateLease {
radix: std::sync::Arc::clone(&self.radix),
namespace: identity.namespace.clone(),
stored_tokens: lookup.stored_tokens.clone(),
};
let token_count = lookup.stored_tokens.len() as u64;
let lookup_ms = lookup_started.elapsed().as_secs_f64() * 1000.0;
let mut reconstruct_ms = 0.0;
let mut reconstruct_bytes = 0u64;
let mut reconstruct_blocks = 0usize;
let mut kv_import_ms = 0.0;
let mut recurrent_import_ms = 0.0;
let mut deterministic_failure = false;
let restore_result = (|| -> Result<bool> {
match lookup.value.payload.kind().into() {
StagePrefixCachePayload::FullState => {
let (full_state, stats) = lookup
.value
.payload
.full_state_bytes_timed()
.context("reconstruct cached full-state payload")
.map_err(|error| {
mark_deterministic_failure(&mut deterministic_failure, error)
})?;
if full_state.is_empty() {
deterministic_failure = true;
return Err(anyhow::anyhow!("cached full-state payload is empty"));
}
add_reconstruct_stats(
&mut reconstruct_ms,
&mut reconstruct_bytes,
&mut reconstruct_blocks,
stats,
);
let import_started = Instant::now();
runtime.import_full_state_for_token_count(
session_id,
full_state.as_ref(),
token_count,
)?;
kv_import_ms = import_started.elapsed().as_secs_f64() * 1000.0;
}
StagePrefixCachePayload::KvRecurrent => {
if let Some((kv, stats)) = lookup
.value
.payload
.kv_bytes_timed()
.context("reconstruct cached KV payload")
.map_err(|error| {
mark_deterministic_failure(&mut deterministic_failure, error)
})?
{
add_reconstruct_stats(
&mut reconstruct_ms,
&mut reconstruct_bytes,
&mut reconstruct_blocks,
stats,
);
if let Some(desc) = lookup.value.extra.kv_desc.as_ref() {
desc.validate_payload(kv.len()).map_err(|error| {
mark_deterministic_failure(&mut deterministic_failure, error)
})?;
if desc.token_start != 0 || desc.token_count != token_count {
deterministic_failure = true;
return Err(anyhow::anyhow!(
"cached KV page token range mismatch for exact-state checkpoint"
));
}
let import_started = Instant::now();
runtime.import_kv_page(session_id, desc, kv.as_ref())?;
kv_import_ms = import_started.elapsed().as_secs_f64() * 1000.0;
} else if !kv.is_empty() {
deterministic_failure = true;
return Err(anyhow::anyhow!(
"cached KV payload is missing its descriptor"
));
}
}
let (recurrent, stats) = lookup
.value
.payload
.recurrent_state_bytes_timed()
.context("reconstruct cached recurrent payload")
.map_err(|error| {
mark_deterministic_failure(&mut deterministic_failure, error)
})?;
if recurrent.is_empty() && !self.dense_without_recurrent {
deterministic_failure = true;
return Err(anyhow::anyhow!("cached recurrent-state payload is empty"));
}
add_reconstruct_stats(
&mut reconstruct_ms,
&mut reconstruct_bytes,
&mut reconstruct_blocks,
stats,
);
let import_started = Instant::now();
if recurrent.is_empty() {
runtime.set_session_position(session_id, token_count)?;
} else {
runtime.import_recurrent_state_for_token_count(
session_id,
recurrent.as_ref(),
token_count,
)?;
}
recurrent_import_ms = import_started.elapsed().as_secs_f64() * 1000.0;
}
_ => return Ok(false),
}
Ok(true)
})();
let restored_payload = match restore_result {
Ok(restored_payload) => restored_payload,
Err(error) => {
drop(lease);
if deterministic_failure
&& let Err(quarantine_error) = self.quarantine_exact_state_entry(
&identity.namespace,
&lookup.stored_tokens,
&lookup.value.page_id,
)
{
let _ =
mesh_llm_events::emit_event(mesh_llm_events::OutputEvent::Warning {
message: "Skippy exact-state quarantine failed".to_string(),
context: Some(format!(
"page_id={} error={quarantine_error:#}",
lookup.value.page_id
)),
});
return Err(error.context(format!(
"failed to fully quarantine corrupt exact-state entry: {quarantine_error:#}"
)));
}
return Err(error);
}
};
if !restored_payload {
drop(lease);
continue;
}
let promote_to_l3 = lookup.value.l3_promotion_eligible
&& self.l3.as_ref().is_some_and(|l3| {
l3.benefit_observe_memory_hit(
&identity.namespace,
&lookup.stored_tokens,
cold_prefill_cost,
)
});
let page_id = lookup.value.page_id.clone();
let promotion_payload = promote_to_l3.then(|| lookup.value.payload.clone());
let promotion_extra = promote_to_l3.then(|| lookup.value.extra.clone());
let stored_tokens = lookup.stored_tokens.clone();
drop(lease);
let promotion_enqueued = if let (Some(payload), Some(extra)) =
(promotion_payload, promotion_extra)
&& let Some(admission_credit) = ExactStateAdmissionCredit::acquire(
&self.admission_outstanding,
&self.admission_best_effort_outstanding,
CaptureAdmission::BestEffort,
)
&& self.try_begin_record(&page_id)
{
matches!(
self.enqueue_exact_state_record(PendingExactStateRecord {
page_id: page_id.clone(),
payload,
extra,
namespace: identity.namespace.clone(),
token_ids: stored_tokens,
l3_fill_claim: None,
write_through_l3: true,
l2_promotion_digest: None,
l3_cost: None,
admission_credit,
}),
ExactStateRecordAdmission::Queued
)
} else {
false
};
let restored = ExactStateRestore {
page_id,
token_count: token_count as usize,
payload_kind: lookup.value.payload.kind(),
logical_bytes: lookup.logical_bytes,
entries,
reconstruct_ms,
reconstruct_bytes,
reconstruct_blocks,
lookup_ms,
kv_import_ms,
recurrent_import_ms,
source: "radix",
fill_ms: 0.0,
rewarm_enqueued: promotion_enqueued,
};
return Ok(Some(restored));
}
Ok(None)
}
fn quarantine_exact_state_entry(
&self,
namespace: &str,
tokens: &[i32],
page_id: &str,
) -> Result<bool> {
let removed = self
.radix
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove_recurrent_if(namespace, tokens, |entry| entry.page_id == page_id);
let Some(entry) = removed else {
return Ok(false);
};
entry.payload.release_from(
&mut self
.exact_blobs
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner),
)?;
Ok(true)
}
pub fn record_exact_state(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identity: &PrefillKvIdentity,
admission: CaptureAdmission,
) -> Result<Option<ExactStateRecord>> {
self.record_exact_state_with_cost(runtime, session_id, identity, admission, None)
}
pub fn record_exact_state_with_cost(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identity: &PrefillKvIdentity,
admission: CaptureAdmission,
l3_cost: Option<skippy_cache::policy::CostSample>,
) -> Result<Option<ExactStateRecord>> {
self.record_exact_state_with_cost_and_durability(
runtime, session_id, identity, admission, l3_cost, true,
)
}
pub(crate) fn record_exact_state_with_cost_and_durability(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identity: &PrefillKvIdentity,
admission: CaptureAdmission,
l3_cost: Option<skippy_cache::policy::CostSample>,
write_through_l3: bool,
) -> Result<Option<ExactStateRecord>> {
let Some(exact_state_payload) = self.exact_state_payload() else {
return Ok(None);
};
if !self.should_record() || !exact_state_payload.is_exact_state() {
#[cfg(test)]
crate::frontend::capture_trace::POST_DECODE_NONE_SHOULD_RECORD
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
return Ok(None);
}
let token_count = identity.identity.token_count;
if token_count < self.checkpoint_policy.min_tokens {
#[cfg(test)]
crate::frontend::capture_trace::POST_DECODE_NONE_MIN_TOKENS
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
return Ok(None);
}
let Some(admission_credit) = ExactStateAdmissionCredit::acquire(
&self.admission_outstanding,
&self.admission_best_effort_outstanding,
admission,
) else {
#[cfg(test)]
{
crate::frontend::capture_trace::POST_DECODE_ADMISSION_DECLINED
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
crate::frontend::capture_trace::ANY_ADMISSION_DECLINED
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
}
return Ok(None);
};
if !self.try_begin_record(&identity.page_id) {
#[cfg(test)]
crate::frontend::capture_trace::POST_DECODE_NONE_BEGIN_RECORD
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
return Ok(None);
}
let already_recorded =
match try_touch_exact_state(&self.radix, &identity.namespace, &identity.token_ids) {
Ok(Some(already_recorded)) => already_recorded,
Ok(None) => {
#[cfg(test)]
crate::frontend::capture_trace::POST_DECODE_NONE_RADIX_BUSY
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
self.finish_record(&identity.page_id);
return Ok(None);
}
Err(error) => {
self.finish_record(&identity.page_id);
return Err(error);
}
};
let probation_recurrence = already_recorded
&& write_through_l3
&& l3_cost.is_some()
&& self.l3.as_ref().is_some_and(|l3| {
l3.benefit_tracks_prefix(&identity.namespace, &identity.token_ids)
});
if already_recorded && !probation_recurrence {
#[cfg(test)]
crate::frontend::capture_trace::POST_DECODE_NONE_ALREADY_RECORDED
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
self.finish_record(&identity.page_id);
return Ok(None);
}
let exported = match exact_state_payload {
StagePrefixCachePayload::FullState => {
runtime.export_full_state(session_id).map(|state| {
(
ExactStatePayload::full_state(state),
ExactStateExtra::default(),
)
})
}
StagePrefixCachePayload::KvRecurrent => (|| {
let kv = match runtime.export_kv_page(session_id, 0, token_count) {
Ok(kv) => Some(kv),
Err(error) if is_native_kv_unavailable(&error) => None,
Err(error) => return Err(error),
};
let recurrent = match runtime.export_recurrent_state(session_id) {
Ok(recurrent) => recurrent,
Err(error)
if self.dense_without_recurrent && is_recurrent_unavailable(&error) =>
{
Vec::new()
}
Err(error) => return Err(error),
};
Ok((
ExactStatePayload::kv_recurrent(
kv.as_ref().map(|kv| kv.payload.clone()).unwrap_or_default(),
recurrent,
),
ExactStateExtra {
kv_desc: kv.as_ref().map(|kv| kv.desc.clone()),
},
))
})(),
StagePrefixCachePayload::Disabled | StagePrefixCachePayload::ResidentKv => {
#[cfg(test)]
crate::frontend::capture_trace::POST_DECODE_NONE_PAYLOAD_DISABLED
.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
self.finish_record(&identity.page_id);
return Ok(None);
}
};
let (payload, extra) = match exported {
Ok(exported) => exported,
Err(error) => {
self.finish_record(&identity.page_id);
return Err(error);
}
};
if payload.byte_len() == 0 {
self.finish_record(&identity.page_id);
return Ok(None);
}
let payload_kind = payload.kind();
let logical_bytes = payload.byte_len();
match self.enqueue_exact_state_record(PendingExactStateRecord {
page_id: identity.page_id.clone(),
payload,
extra,
namespace: identity.namespace.clone(),
token_ids: identity.token_ids.clone(),
l3_fill_claim: None,
write_through_l3,
l2_promotion_digest: None,
l3_cost,
admission_credit,
}) {
ExactStateRecordAdmission::Queued => {
let entries = self
.radix
.try_lock()
.ok()
.map(|radix| radix.stats().recurrent_entries)
.unwrap_or_default();
let physical_bytes = self
.exact_blobs
.try_lock()
.ok()
.map(|blobs| blobs.physical_bytes())
.unwrap_or_default();
Ok(Some(ExactStateRecord {
page_id: identity.page_id.clone(),
token_count: token_count as usize,
payload_kind,
stored: false,
logical_bytes,
physical_bytes,
entries,
evicted_entries: 0,
evicted_logical_bytes: 0,
dedupe: Default::default(),
}))
}
ExactStateRecordAdmission::DroppedFull | ExactStateRecordAdmission::WorkerStopped => {
Ok(None)
}
}
}
}
impl KvStageIntegration {
fn restore_from_l3(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identity: &PrefillKvIdentity,
lookup_started: Instant,
cold_prefill_cost: Option<f64>,
shared_durable_fallback: bool,
) -> Result<Option<ExactStateRestore>> {
let Some(l3) = &self.l3 else {
return Ok(None);
};
let durable_token_ids = if shared_durable_fallback {
self.durable_exact_lookup_token_ids(&identity.token_ids)
} else {
&identity.token_ids
};
let location =
match l3.locate_longest(&identity.namespace, durable_token_ids, MAX_PREFIX_PROBES) {
Ok(Some(location)) => location,
Ok(None) | Err(_) => return Ok(None),
};
if let Some(restored) =
self.restore_from_l2(runtime, session_id, identity, &location, lookup_started)?
{
return Ok(Some(restored));
}
let fill_claim = l3_fill_claim_key(l3, &location);
{
let mut inflight = self
.inflight_fills
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if !inflight.insert(fill_claim.clone()) {
return Ok(None);
}
}
let outcome = self.fill_and_import(
runtime,
session_id,
identity,
lookup_started,
l3,
&location,
cold_prefill_cost,
);
let handed_to_worker = matches!(&outcome, Ok(Some(restored)) if restored.rewarm_enqueued);
if !handed_to_worker {
self.inflight_fills
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(&fill_claim);
}
outcome
}
#[allow(clippy::too_many_arguments)]
fn fill_and_import(
&self,
runtime: &mut RuntimeState,
session_id: &str,
identity: &PrefillKvIdentity,
lookup_started: Instant,
l3: &std::sync::Arc<skippy_cache::L3Tier>,
location: &skippy_cache::L3Location,
cold_prefill_cost: Option<f64>,
) -> Result<Option<ExactStateRestore>> {
let Ok(native_kv_desc) = preflight_l3_kv_location(location) else {
return Ok(None);
};
let fill_started = Instant::now();
let fill = match l3.load(location) {
Ok(fill) => fill,
Err(_) => {
if let Some(l2) = &self.l2 {
l2.invalidate_digest(&location.manifest_key);
}
return Ok(None);
}
};
if fill.payload.byte_len() == 0 {
return Ok(None);
}
if (location.native_kv_passthrough || location.cachegen_kv)
&& (fill.token_count != location.token_count
|| fill.kv_desc_json != location.kv_desc_json
|| fill.cachegen_kv != location.cachegen_kv
|| fill.kv_decoded_bytes != location.kv_decoded_bytes)
{
return Ok(None);
}
let fill_ms = fill_started.elapsed().as_secs_f64() * 1000.0;
let token_count = fill.token_count;
let kv_desc: Option<skippy_runtime::RuntimeKvPageDesc> = native_kv_desc.or_else(|| {
fill.kv_desc_json
.as_deref()
.and_then(|json| serde_json::from_str(json).ok())
});
let lookup_ms = lookup_started.elapsed().as_secs_f64() * 1000.0;
let mut kv_import_ms = 0.0;
let mut recurrent_import_ms = 0.0;
match fill.payload.kind().into() {
StagePrefixCachePayload::FullState => {
let (full_state, _) = fill
.payload
.full_state_bytes_timed()
.context("reconstruct L3 full-state payload")?;
if full_state.is_empty() {
return Ok(None);
}
let import_started = Instant::now();
runtime.import_full_state_for_token_count(
session_id,
full_state.as_ref(),
token_count,
)?;
kv_import_ms = import_started.elapsed().as_secs_f64() * 1000.0;
}
StagePrefixCachePayload::KvRecurrent => {
let kv = fill
.payload
.kv_bytes()
.context("reconstruct L3 KV payload")?;
let recurrent = fill
.payload
.recurrent_state_bytes()
.context("reconstruct L3 recurrent payload")?;
if recurrent.is_empty() && !self.dense_without_recurrent {
return Ok(None);
}
let kv_page = match (kv.as_ref(), kv_desc.as_ref()) {
(Some(kv), Some(desc)) => {
let payload_valid = if fill.cachegen_kv {
usize::try_from(desc.payload_bytes)
.ok()
.is_some_and(|raw_len| {
skippy_cache::cachegen::archive::validate_archive(kv, raw_len)
.is_ok()
})
} else {
desc.validate_payload(kv.len()).is_ok()
};
if !payload_valid
|| desc.token_start != 0
|| desc.token_count != token_count
{
return Ok(None);
}
Some((kv, desc))
}
(Some(kv), None) if !kv.is_empty() => return Ok(None),
_ => None,
};
if let Some((kv, desc)) = kv_page {
let import_started = Instant::now();
if fill.cachegen_kv {
runtime.import_cachegen_kv_page(session_id, desc, kv.as_ref())?;
} else {
runtime.import_kv_page(session_id, desc, kv.as_ref())?;
}
kv_import_ms = import_started.elapsed().as_secs_f64() * 1000.0;
}
let import_started = Instant::now();
if recurrent.is_empty() {
runtime.set_session_position(session_id, token_count)?;
} else {
runtime.import_recurrent_state_for_token_count(
session_id,
recurrent.as_ref(),
token_count,
)?;
}
recurrent_import_ms = import_started.elapsed().as_secs_f64() * 1000.0;
}
_ => return Ok(None),
}
let logical_bytes = if fill.cachegen_kv {
fill.kv_decoded_bytes.saturating_add(
fill.payload
.recurrent_state_bytes()
.map(|bytes| bytes.len() as u64)
.unwrap_or_default(),
)
} else {
fill.payload.byte_len()
};
let payload_kind = fill.payload.kind();
let restore_cost = lookup_started.elapsed().as_secs_f64() * 1_000.0;
l3.benefit_observe_l3_restore(
&identity.namespace,
&identity.token_ids[..token_count as usize],
location,
cold_prefill_cost,
restore_cost,
);
if fill.cachegen_kv {
return Ok(Some(ExactStateRestore {
page_id: identity.page_id.clone(),
token_count: token_count as usize,
payload_kind,
logical_bytes,
entries: 0,
reconstruct_ms: 0.0,
reconstruct_bytes: 0,
reconstruct_blocks: 0,
lookup_ms,
kv_import_ms,
recurrent_import_ms,
source: "l3-cachegen",
fill_ms,
rewarm_enqueued: false,
}));
}
let l2_promotion_digest = self.l2.as_ref().and_then(|l2| {
l2.consider_l3_fill(&location.manifest_key, token_count, fill.payload.byte_len())
.then(|| location.manifest_key.clone())
});
let rewarm_enqueued = ExactStateAdmissionCredit::acquire(
&self.admission_outstanding,
&self.admission_best_effort_outstanding,
CaptureAdmission::BestEffort,
)
.is_some_and(|admission_credit| {
matches!(
self.enqueue_exact_state_record(PendingExactStateRecord {
page_id: identity.page_id.clone(),
payload: fill.payload,
extra: ExactStateExtra { kv_desc },
namespace: identity.namespace.clone(),
token_ids: identity.token_ids[..token_count as usize].to_vec(),
l3_fill_claim: Some(l3_fill_claim_key(l3, location)),
write_through_l3: false,
l2_promotion_digest,
l3_cost: None,
admission_credit,
}),
ExactStateRecordAdmission::Queued
)
});
Ok(Some(ExactStateRestore {
page_id: identity.page_id.clone(),
token_count: token_count as usize,
payload_kind,
logical_bytes,
entries: 0,
reconstruct_ms: 0.0,
reconstruct_bytes: 0,
reconstruct_blocks: 0,
lookup_ms,
kv_import_ms,
recurrent_import_ms,
source: "l3",
fill_ms,
rewarm_enqueued,
}))
}
}
fn is_recurrent_unavailable(error: &anyhow::Error) -> bool {
error
.chain()
.any(|cause| cause.to_string().contains("no recurrent memory"))
}
struct ExactStateLease {
radix: std::sync::Arc<
std::sync::Mutex<
skippy_cache::UnifiedRadixCache<super::RadixResidentEntry, super::RadixExactEntry>,
>,
>,
namespace: String,
stored_tokens: Vec<i32>,
}
impl Drop for ExactStateLease {
fn drop(&mut self) {
let released = self
.radix
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.release_recurrent(&self.namespace, &self.stored_tokens);
debug_assert!(released, "recurrent radix acquire/release must balance");
}
}
fn try_touch_exact_state(
radix: &std::sync::Mutex<
skippy_cache::UnifiedRadixCache<super::RadixResidentEntry, super::RadixExactEntry>,
>,
namespace: &str,
token_ids: &[i32],
) -> Result<Option<bool>> {
match radix.try_lock() {
Ok(mut radix) => Ok(Some(radix.recurrent_exact(namespace, token_ids).is_some())),
Err(std::sync::TryLockError::WouldBlock) => Ok(None),
Err(std::sync::TryLockError::Poisoned(poisoned)) => Ok(Some(
poisoned
.into_inner()
.recurrent_exact(namespace, token_ids)
.is_some(),
)),
}
}
fn is_native_kv_unavailable(error: &anyhow::Error) -> bool {
error.chain().any(|cause| {
let message = cause.to_string();
message.contains("runtime memory type is not supported for native KV pages")
|| message.contains("runtime has no attention KV cache")
|| message.contains("no KV cache layers selected by layer range")
})
}
fn mark_deterministic_failure(
deterministic_failure: &mut bool,
error: anyhow::Error,
) -> anyhow::Error {
*deterministic_failure = true;
error
}
impl StagePrefixCachePayload {
pub(crate) fn is_exact_state(self) -> bool {
matches!(self, Self::KvRecurrent | Self::FullState)
}
}
impl From<skippy_cache::ExactStatePayloadKind> for StagePrefixCachePayload {
fn from(kind: skippy_cache::ExactStatePayloadKind) -> Self {
match kind {
skippy_cache::ExactStatePayloadKind::FullState => Self::FullState,
skippy_cache::ExactStatePayloadKind::KvRecurrent => Self::KvRecurrent,
skippy_cache::ExactStatePayloadKind::RecurrentOnly => Self::Disabled,
}
}
}
impl From<StagePrefixCachePayload> for skippy_cache::ExactStatePayloadKind {
fn from(payload: StagePrefixCachePayload) -> Self {
match payload {
StagePrefixCachePayload::FullState => Self::FullState,
StagePrefixCachePayload::KvRecurrent => Self::KvRecurrent,
StagePrefixCachePayload::Disabled | StagePrefixCachePayload::ResidentKv => {
Self::FullState
}
}
}
}
#[cfg(test)]
mod tests {
use std::{
sync::{Arc, Mutex},
time::{Duration, Instant},
};
use skippy_cache::{L3Location, UnifiedRadixCache};
use super::{
is_native_kv_unavailable, preflight_l3_kv_location, resident_prefix_is_preferred,
try_touch_exact_state,
};
type TestRadix = UnifiedRadixCache<
crate::kv_integration::RadixResidentEntry,
crate::kv_integration::RadixExactEntry,
>;
#[test]
fn longer_resident_prefixes_skip_exact_restore() {
assert!(resident_prefix_is_preferred(4_000, None));
assert!(resident_prefix_is_preferred(4_001, Some(4_000)));
assert!(!resident_prefix_is_preferred(4_000, Some(4_000)));
assert!(!resident_prefix_is_preferred(200, Some(4_000)));
}
#[test]
fn empty_kv_layer_range_is_an_unavailable_optional_kv_component() {
let error = anyhow::anyhow!("RuntimeError: no KV cache layers selected by layer range");
assert!(is_native_kv_unavailable(&error));
}
#[test]
fn busy_exact_state_lock_skips_touch_without_waiting() {
let cache = Arc::new(Mutex::new(TestRadix::new()));
let locked = cache.clone();
let (locked_tx, locked_rx) = std::sync::mpsc::channel();
let (release_tx, release_rx) = std::sync::mpsc::channel();
let holder = std::thread::spawn(move || {
let _guard = locked.lock().unwrap();
locked_tx.send(()).unwrap();
release_rx.recv().unwrap();
});
locked_rx.recv().unwrap();
let started = Instant::now();
assert_eq!(
try_touch_exact_state(&cache, "namespace", &[1]).unwrap(),
None
);
assert!(started.elapsed() < Duration::from_millis(100));
release_tx.send(()).unwrap();
holder.join().unwrap();
}
#[test]
fn poisoned_exact_state_lock_recovers_without_panicking() {
let cache = Arc::new(Mutex::new(TestRadix::new()));
let poisoned = cache.clone();
assert!(
std::thread::spawn(move || {
let _guard = poisoned.lock().unwrap();
panic!("poison exact-state cache for test");
})
.join()
.is_err()
);
assert_eq!(
try_touch_exact_state(&cache, "namespace", &[1]).unwrap(),
Some(false)
);
}
fn native_location(desc: &skippy_runtime::RuntimeKvPageDesc) -> L3Location {
L3Location {
namespace_key: "namespace".to_string(),
prefix_key: "prefix".to_string(),
token_count: desc.token_count,
manifest_key: "manifest".to_string(),
kv_desc_json: Some(serde_json::to_string(desc).unwrap()),
kv_bytes: desc.payload_bytes,
native_kv_passthrough: true,
cachegen_kv: false,
kv_decoded_bytes: desc.payload_bytes,
}
}
fn native_desc() -> skippy_runtime::RuntimeKvPageDesc {
skippy_runtime::RuntimeKvPageDesc {
version: 1,
layer_start: 0,
layer_end: 1,
token_start: 0,
token_count: 8,
layer_count: 1,
k_type: skippy_runtime::GGML_TYPE_Q8_0,
v_type: skippy_runtime::GGML_TYPE_Q8_0,
k_row_bytes: 16,
v_row_bytes: 16,
v_element_bytes: 2,
k_idx_row_bytes: 0,
payload_bytes: 256,
flags: 0,
codec: 0,
component_count: 0,
components: Box::new([Default::default(); 2]),
}
}
#[test]
fn native_kv_descriptor_is_validated_before_segment_load() {
let desc = native_desc();
assert_eq!(
preflight_l3_kv_location(&native_location(&desc)).unwrap(),
Some(desc)
);
let mut wrong_length = native_desc();
wrong_length.payload_bytes += 1;
let mut location = native_location(&wrong_length);
location.kv_bytes -= 1;
assert!(preflight_l3_kv_location(&location).is_err());
let mut wrong_prefix = native_desc();
wrong_prefix.token_start = 1;
assert!(preflight_l3_kv_location(&native_location(&wrong_prefix)).is_err());
}
}
#[cfg(test)]
mod admission_credit_tests {
use super::{
CaptureAdmission, EXACT_STATE_ADMISSION_RESERVED, EXACT_STATE_ADMISSION_TOTAL,
ExactStateAdmissionCredit,
};
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
fn counters() -> (Arc<AtomicUsize>, Arc<AtomicUsize>) {
(Arc::new(AtomicUsize::new(0)), Arc::new(AtomicUsize::new(0)))
}
#[test]
fn construction_and_drop_account_exactly_once() {
let (total, best_effort) = counters();
let guard = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
);
assert!(guard.is_some(), "fresh budget admits Continuation");
assert_eq!(total.load(Ordering::Acquire), 1);
drop(guard);
assert_eq!(total.load(Ordering::Acquire), 0);
}
#[test]
fn continuation_declines_when_total_budget_exhausted() {
let (total, best_effort) = counters();
let first = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
);
let second = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
);
assert!(first.is_some() && second.is_some(), "two slots for TOTAL=2");
let third = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
);
assert!(
third.is_none(),
"third Continuation must decline at TOTAL=2"
);
assert_eq!(total.load(Ordering::Acquire), 2);
}
#[test]
fn best_effort_declines_at_reserved_cap_even_with_budget_headroom() {
let (total, best_effort) = counters();
let first =
ExactStateAdmissionCredit::acquire(&total, &best_effort, CaptureAdmission::BestEffort);
assert!(first.is_some());
let second =
ExactStateAdmissionCredit::acquire(&total, &best_effort, CaptureAdmission::BestEffort);
assert!(
second.is_none(),
"BestEffort cap is TOTAL - RESERVED = 1; second must decline"
);
assert_eq!(best_effort.load(Ordering::Acquire), 1);
}
#[test]
fn continuation_admits_after_best_effort_under_reserved_share() {
let (total, best_effort) = counters();
let best =
ExactStateAdmissionCredit::acquire(&total, &best_effort, CaptureAdmission::BestEffort);
assert!(best.is_some(), "BestEffort takes its single slot");
let cont = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
);
assert!(
cont.is_some(),
"Continuation must admit into the reserved share"
);
}
#[test]
fn concurrent_producers_never_oversubscribe_the_budget() {
use std::sync::Barrier;
const THREADS: usize = 4;
const ATTEMPTS: usize = 25;
let (total, best_effort) = counters();
let start = Arc::new(Barrier::new(THREADS));
let handles: Vec<_> = (0..THREADS)
.map(|thread| {
let total = Arc::clone(&total);
let best_effort = Arc::clone(&best_effort);
let start = Arc::clone(&start);
std::thread::spawn(move || {
let mut held = Vec::new();
start.wait();
for attempt in 0..ATTEMPTS {
let class = if (thread + attempt) % 2 == 0 {
CaptureAdmission::Continuation
} else {
CaptureAdmission::BestEffort
};
if let Some(credit) =
ExactStateAdmissionCredit::acquire(&total, &best_effort, class)
{
held.push((class, credit));
}
}
held
})
})
.collect();
let mut held: Vec<_> = handles
.into_iter()
.flat_map(|handle| handle.join().unwrap())
.collect();
let held_best_effort = held
.iter()
.filter(|(class, _)| *class == CaptureAdmission::BestEffort)
.count();
assert!(held.len() <= EXACT_STATE_ADMISSION_TOTAL);
assert!(
held_best_effort <= EXACT_STATE_ADMISSION_TOTAL - EXACT_STATE_ADMISSION_RESERVED,
"BestEffort holds must respect the reserved share"
);
assert_eq!(
total.load(Ordering::Acquire),
held.len(),
"retained credits must account for every outstanding slot"
);
assert_eq!(
best_effort.load(Ordering::Acquire),
held_best_effort,
"retained credits must account for the BestEffort counter"
);
while held.len() < EXACT_STATE_ADMISSION_TOTAL {
let credit = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
)
.expect("topping up an under-subscribed budget must succeed");
held.push((CaptureAdmission::Continuation, credit));
}
let hammer = Arc::new(Barrier::new(THREADS));
let hammers: Vec<_> = (0..THREADS)
.map(|_| {
let total = Arc::clone(&total);
let best_effort = Arc::clone(&best_effort);
let hammer = Arc::clone(&hammer);
std::thread::spawn(move || {
hammer.wait();
for _ in 0..ATTEMPTS {
assert!(
ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::BestEffort
)
.is_none(),
"an exhausted budget must decline every BestEffort attempt"
);
}
})
})
.collect();
for handle in hammers {
handle.join().unwrap();
}
assert_eq!(
total.load(Ordering::Acquire),
held.len(),
"the rollback hammer must not change the total"
);
assert_eq!(
best_effort.load(Ordering::Acquire),
held_best_effort,
"the rollback hammer must not leak a BestEffort slot"
);
drop(held);
assert_eq!(total.load(Ordering::Acquire), 0);
assert_eq!(best_effort.load(Ordering::Acquire), 0);
let recovery_continuation = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
);
let recovery_best_effort =
ExactStateAdmissionCredit::acquire(&total, &best_effort, CaptureAdmission::BestEffort);
assert!(
recovery_continuation.is_some() && recovery_best_effort.is_some(),
"budget must recover after the drain"
);
}
#[test]
fn rollback_is_exercised_when_the_budget_is_held_by_continuations() {
use std::sync::Barrier;
const THREADS: usize = 4;
const ATTEMPTS: usize = 50;
let (total, best_effort) = counters();
let first = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
)
.expect("first Continuation admitted");
let second = ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::Continuation,
)
.expect("second Continuation admitted");
assert_eq!(total.load(Ordering::Acquire), 2);
assert_eq!(best_effort.load(Ordering::Acquire), 0);
let barrier = Arc::new(Barrier::new(THREADS));
let hammers: Vec<_> = (0..THREADS)
.map(|_| {
let total = Arc::clone(&total);
let best_effort = Arc::clone(&best_effort);
let barrier = Arc::clone(&barrier);
std::thread::spawn(move || {
barrier.wait();
for _ in 0..ATTEMPTS {
assert!(
ExactStateAdmissionCredit::acquire(
&total,
&best_effort,
CaptureAdmission::BestEffort
)
.is_none(),
"a continuation-exhausted budget must decline every BestEffort via rollback"
);
}
})
})
.collect();
for handle in hammers {
handle.join().unwrap();
}
assert_eq!(
total.load(Ordering::Acquire),
2,
"rollback leaves the total at the held credits"
);
assert_eq!(
best_effort.load(Ordering::Acquire),
0,
"rollback must not leak a BestEffort slot"
);
drop((first, second));
assert_eq!(total.load(Ordering::Acquire), 0);
assert_eq!(best_effort.load(Ordering::Acquire), 0);
assert!(
ExactStateAdmissionCredit::acquire(&total, &best_effort, CaptureAdmission::BestEffort)
.is_some(),
"budget recovers after the drain"
);
}
}