//! A single atomic checkpoint for the existing DS4 KV index, not a second bank. //! Tensor payloads keep the reference codec's dtypes, shapes and raw bytes. use super::request::LiveSession; use super::session_codec::decode_payload; use super::stream::{Stream, Streams}; use crate::engine::metal::checkpoint::{read_u64, write_u64}; use serde::{Deserialize, Serialize}; use std::collections::BTreeMap; use std::fs::{self, File}; use std::io::{Read, Seek, SeekFrom, Write}; use std::path::Path; const MAGIC: &[u8; 8] = b"QWMTKV01"; const HEADER_BYTES: u64 = 16; const MAX_MANIFEST_BYTES: u64 = 64 * 1024 * 1024; const IO_CHUNK: usize = 8 * 1024 * 1024; #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] pub(super) struct Identity { pub(super) model: [u8; 32], pub(super) context: u32, pub(super) vocabulary: u32, pub(super) mtp: bool, } #[derive(Serialize, Deserialize)] struct Manifest { identity: Identity, tag: [u8; 32], tokens: Vec, lazy_kv: bool, payload: serde_json::Value, // Offsets are absolute; blob names are codec-local, never filesystem paths. blobs: BTreeMap, } fn check_abort(abort: &impl Fn() -> bool) -> Result<(), String> { if abort() { Err("Qwen checkpoint interrupted".into()) } else { Ok(()) } } fn validate_tokens(tokens: &[i32], identity: Identity) -> Result<(), String> { if identity.context == 0 || identity.vocabulary == 0 || tokens.is_empty() || tokens.len() > identity.context as usize || tokens .iter() .any(|&token| token < 0 || token as u32 >= identity.vocabulary) { return Err("Qwen checkpoint token prefix is invalid".into()); } Ok(()) } impl LiveSession { #[allow(clippy::too_many_arguments)] pub(super) fn save_checkpoint( &mut self, path: &Path, identity: Identity, tag: [u8; 32], streams: &Streams, stream: Stream, mut progress: impl FnMut(u64), abort: impl Fn() -> bool, ) -> Result<(), String> { check_abort(&abort)?; let saved = self .current .as_ref() .ok_or("Qwen has no committed checkpoint")?; validate_tokens(&saved.token_ids, identity)?; if saved.mtp.is_some() != identity.mtp { return Err("Qwen checkpoint MTP configuration differs".into()); } if let Some(parent) = path.parent() { fs::create_dir_all(parent).map_err(|e| e.to_string())?; } // Same serialized-writer, staging-file and sync/rename contract as DS4. let temporary = path.with_extension("tmp"); let result = (|| -> Result<(), String> { let mut file = File::create(&temporary).map_err(|e| e.to_string())?; file.write_all(MAGIC).map_err(|e| e.to_string())?; write_u64(&mut file, 0)?; progress(HEADER_BYTES); let mut blobs = BTreeMap::new(); let payload = saved.encode_payload( 256, streams, stream, |name, bytes| { let offset = file.stream_position().map_err(|e| e.to_string())?; for chunk in bytes.chunks(IO_CHUNK) { check_abort(&abort)?; file.write_all(chunk).map_err(|e| e.to_string())?; progress(chunk.len() as u64); } if blobs.insert(name, (offset, bytes.len() as u64)).is_some() { return Err("duplicate Qwen checkpoint blob".into()); } Ok(()) }, &abort, )?; check_abort(&abort)?; let offset = file.stream_position().map_err(|e| e.to_string())?; let manifest = serde_json::to_vec(&Manifest { identity, tag, tokens: saved.token_ids.clone(), lazy_kv: saved.lazy_kv, payload, blobs, }) .map_err(|e| e.to_string())?; if manifest.len() as u64 > MAX_MANIFEST_BYTES { return Err("Qwen checkpoint metadata exceeds limit".into()); } file.write_all(&manifest).map_err(|e| e.to_string())?; progress(manifest.len() as u64); file.seek(SeekFrom::Start(8)).map_err(|e| e.to_string())?; write_u64(&mut file, offset)?; file.sync_all().map_err(|e| e.to_string())?; check_abort(&abort)?; fs::rename(&temporary, path).map_err(|e| e.to_string()) })(); if result.is_err() { let _ = fs::remove_file(&temporary); } else { self.tag = tag; } result } #[allow(clippy::too_many_arguments)] pub(super) fn load_checkpoint( &mut self, path: &Path, identity: Identity, streams: &Streams, stream: Stream, mut progress: impl FnMut(u64), abort: impl Fn() -> bool, ) -> Result { check_abort(&abort)?; let mut file = match File::open(path) { Ok(file) => file, Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(false), Err(e) => return Err(e.to_string()), }; let mut magic = [0; 8]; file.read_exact(&mut magic).map_err(|e| e.to_string())?; if &magic != MAGIC { return Err("Qwen checkpoint format differs; rebuild the cache".into()); } let offset = read_u64(&mut file)?; let length = file.metadata().map_err(|e| e.to_string())?.len(); if offset < HEADER_BYTES || offset >= length || length - offset > MAX_MANIFEST_BYTES { return Err("Qwen checkpoint metadata range is invalid".into()); } file.seek(SeekFrom::Start(offset)) .map_err(|e| e.to_string())?; let mut raw = vec![0; (length - offset) as usize]; file.read_exact(&mut raw).map_err(|e| e.to_string())?; let mut manifest: Manifest = serde_json::from_slice(&raw).map_err(|e| e.to_string())?; progress(HEADER_BYTES + raw.len() as u64); if manifest.identity != identity { return Err( "Qwen checkpoint does not match the current model/context/MTP policy".into(), ); } validate_tokens(&manifest.tokens, identity)?; if !manifest.payload["mtp_history_snapshot"].is_null() != identity.mtp { return Err("Qwen checkpoint MTP payload differs".into()); } // Validate all extents before allocating any GPU state. Require exact // coverage, so overlapping blobs and unindexed/trailing bytes fail closed. let mut ranges = manifest.blobs.values().copied().collect::>(); ranges.sort_unstable(); let mut end = HEADER_BYTES; for (start, size) in ranges { if start != end { return Err("Qwen checkpoint blob range is invalid".into()); } end = start .checked_add(size) .filter(|&n| n <= offset) .ok_or("Qwen checkpoint blob exceeds payload")?; } if end != offset { return Err("Qwen checkpoint payload coverage differs".into()); } let saved = decode_payload( &manifest.payload, manifest.tokens, manifest.lazy_kv, true, streams, stream, |name| { check_abort(&abort)?; let (start, size) = manifest .blobs .remove(name) .ok_or("missing or repeated Qwen checkpoint blob")?; file.seek(SeekFrom::Start(start)) .map_err(|e| e.to_string())?; let mut bytes = Vec::new(); let size = usize::try_from(size).map_err(|_| "Qwen checkpoint blob too large")?; bytes.try_reserve_exact(size).map_err(|e| e.to_string())?; bytes.resize(size, 0); for chunk in bytes.chunks_mut(IO_CHUNK) { check_abort(&abort)?; file.read_exact(chunk).map_err(|e| e.to_string())?; progress(chunk.len() as u64); } Ok(bytes) }, )?; if !manifest.blobs.is_empty() { return Err("unreferenced Qwen checkpoint blobs".into()); } check_abort(&abort)?; // Publish only after complete validation and materialization. Failure // leaves the live conversation and its prefix untouched. self.current = Some(saved); self.prefix = None; self.tag = manifest.tag; Ok(true) } } #[test] #[ignore = "model-free Metal checkpoint; external supervisor and DS4_QWEN38_CHECKPOINT_DIR required"] fn qwen_checkpoint_roundtrip_and_failed_restore_preserve_live_state() { use super::super::gpu::Context; use super::session_cache::{Boundary, SessionSnapshot, State}; use super::stream::Device; use super::{ allocator::Allocator, array::{Array, Dtype}, configure_sources, scalar_buffer, }; use std::cell::Cell; configure_sources().unwrap(); let _context = Context::open_qwen(0).unwrap(); let streams = Streams::new(Some(Allocator::new().unwrap())); let stream = streams.default_stream(Device::Gpu).unwrap(); let leaf = Array::new( &[1, 1, 4], Dtype::BF16, scalar_buffer(&[0, 0, 128, 63, 0, 64, 64, 64]).unwrap(), ) .unwrap(); let state = Some(State::Recurrent(vec![Some(leaf.clone()), None])); let mut live = LiveSession { current: Some(SessionSnapshot { token_ids: vec![1, 2, 3, 4], logits: leaf.clone(), hidden: Some(leaf.clone()), trunk: vec![state.clone()], mtp: Some(vec![None]), lazy_kv: true, boundaries: vec![Boundary { tokens: 2, state: vec![state], hidden: Some(leaf), }], }), ..Default::default() }; let identity = Identity { model: [7; 32], context: 64, vocabulary: 256, mtp: true, }; let directory = std::path::PathBuf::from(std::env::var_os("DS4_QWEN38_CHECKPOINT_DIR").unwrap()); fs::create_dir_all(&directory).unwrap(); let path = directory.join("roundtrip.bin"); let fingerprint = |live: &LiveSession| { let mut blobs = BTreeMap::new(); let payload = live .current .as_ref() .unwrap() .encode_payload( 256, &streams, stream, |name, bytes| { blobs.insert(name, bytes); Ok(()) }, || false, ) .unwrap(); (payload, blobs) }; let original = fingerprint(&live); let mut written = 0; live.save_checkpoint( &path, identity, [9; 32], &streams, stream, |n| written += n, || false, ) .unwrap(); let valid = fs::read(&path).unwrap(); assert_eq!(written, valid.len() as u64); assert_eq!(live.tag, [9; 32]); live = LiveSession::default(); let mut read = 0; assert!( live.load_checkpoint(&path, identity, &streams, stream, |n| read += n, || false) .unwrap() ); assert_eq!(read, written); assert_eq!(fingerprint(&live), original); assert_eq!(live.tag, [9; 32]); assert_eq!(live.current.as_ref().unwrap().token_ids, [1, 2, 3, 4]); assert!(live.current.as_ref().unwrap().lazy_kv); assert!( !live .load_checkpoint( &directory.join("missing.bin"), identity, &streams, stream, |_| {}, || false ) .unwrap() ); // Cancellation after IO has started must not replace a valid file/tag. let interrupted = Cell::new(false); assert!( live.save_checkpoint( &path, identity, [8; 32], &streams, stream, |_| interrupted.set(true), || interrupted.get() ) .is_err() ); assert_eq!(fs::read(&path).unwrap(), valid); assert!(!path.with_extension("tmp").exists()); interrupted.set(false); assert!( live.load_checkpoint( &path, identity, &streams, stream, |_| interrupted.set(true), || interrupted.get() ) .is_err() ); for changed in [ Identity { context: 65, ..identity }, Identity { mtp: false, ..identity }, Identity { model: [8; 32], ..identity }, ] { assert!( live.load_checkpoint(&path, changed, &streams, stream, |_| {}, || false) .is_err() ); } let offset = u64::from_le_bytes(valid[8..16].try_into().unwrap()) as usize; let broken = directory.join("broken.bin"); let mutations: [fn(&mut Manifest); 5] = [ |m| m.tokens[0] = -1, |m| m.tokens.resize(65, 1), |m| m.blobs.values_mut().next().unwrap().1 = u64::MAX, |m| m.blobs.values_mut().next().unwrap().0 = 0, |m| m.payload["logits"]["nbytes"] = serde_json::json!(99), ]; for mutate in mutations { let mut manifest: Manifest = serde_json::from_slice(&valid[offset..]).unwrap(); mutate(&mut manifest); let mut bytes = valid[..offset].to_vec(); bytes.extend(serde_json::to_vec(&manifest).unwrap()); fs::write(&broken, bytes).unwrap(); assert!( live.load_checkpoint(&broken, identity, &streams, stream, |_| {}, || false) .is_err() ); } for bytes in [&b"DS4QWN01"[..], &valid[..offset], &valid[..7]] { fs::write(&broken, bytes).unwrap(); assert!( live.load_checkpoint(&broken, identity, &streams, stream, |_| {}, || false) .is_err() ); } assert_eq!(live.tag, [9; 32]); assert_eq!(fingerprint(&live), original); println!( "Qwen checkpoint: exact payload roundtrip, interrupted save, incompatible/corrupt restore passed" ); }