Commit 033df86b6 for llama.cpp
commit 033df86b69ec1a333eb241f0c16325a4d43dcff5
Author: Leebr Data Consulting <harrak.amine1987@gmail.com>
Date: Thu Oct 8 12:18:20 2026 +0200
server : preserve context checkpoints across slot save/restore (#26004)
* server : preserve context checkpoints across slot save/restore
Append the checkpoints after the packed server_tokens payload added in #26640
and count them in n_written / n_read, so a restored slot can still roll back to
a checkpoint instead of re-processing the whole prompt.
* server : drop draft checkpoint data that does not match the draft context
Restoring a slot saved with a different draft KV cache type aborted in
load_dft(). Test-load one draft checkpoint on restore and drop the draft
data if it does not fit, instead of crashing. Adds a regression test.
Co-authored-by: Igor Okulist <okigan@gmail.com>
* server : harden the checkpoint appendix of slot save files
Bound each blob size by the bytes left in the file before allocating, open the
file with UTF-8 paths on Windows like the llama state payload, fall back to full
prompt re-processing when a checkpoint restored from a slot file fails to load,
and replace the 1024 count cap by keeping the last n_ctx_checkpoints while reading.
* server : report an incomplete checkpoint appendix as a failed slot save
Return an error to the client when the appendix cannot be written, like a
failed payload write, and make the oversized-blob test declare a size that
cannot be allocated, so an unbounded allocation fails the test.
* server : reject an empty target state in the checkpoint appendix
A saved checkpoint always holds a target state, an empty blob would roll back
without restoring anything. Also log with the slot id, and load the draft test
model from the HF cache instead of a second download.
* common : return bool from checkpoint load_tgt / load_dft
A checkpoint restored from a slot file falls back to full prompt re-processing
when it fails to load, a checkpoint created in memory still aborts.
---------
Co-authored-by: Igor Okulist <okigan@gmail.com>
diff --git a/common/common.cpp b/common/common.cpp
index 28ea8680e..48caa80c1 100644
--- a/common/common.cpp
+++ b/common/common.cpp
@@ -2388,40 +2388,36 @@ void common_prompt_checkpoint::update_dft(
}
}
-void common_prompt_checkpoint::load_tgt(
+bool common_prompt_checkpoint::load_tgt(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const {
if (ctx == nullptr) {
- return;
+ return true;
}
if (data_tgt.empty()) {
- return;
+ return true;
}
const size_t n = llama_state_seq_set_data_ext(ctx, data_tgt.data(), data_tgt.size(), seq_id, flags);
- if (n != data_tgt.size()) {
- GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_tgt.size(), n);
- }
+ return n == data_tgt.size();
}
-void common_prompt_checkpoint::load_dft(
+bool common_prompt_checkpoint::load_dft(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const {
if (ctx == nullptr) {
- return;
+ return true;
}
if (data_dft.empty()) {
- return;
+ return true;
}
const size_t n = llama_state_seq_set_data_ext(ctx, data_dft.data(), data_dft.size(), seq_id, flags);
- if (n != data_dft.size()) {
- GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_dft.size(), n);
- }
+ return n == data_dft.size();
}
void common_prompt_checkpoint::clear_tgt() {
diff --git a/common/common.h b/common/common.h
index 0a85f11f9..de88dfb9b 100644
--- a/common/common.h
+++ b/common/common.h
@@ -1296,12 +1296,13 @@ struct common_prompt_checkpoint {
llama_seq_id seq_id,
llama_state_seq_flags flags);
- void load_tgt(
+ // return false if the state could not be restored
+ bool load_tgt(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const;
- void load_dft(
+ bool load_dft(
llama_context * ctx,
llama_seq_id seq_id,
llama_state_seq_flags flags) const;
diff --git a/examples/speculative-simple/speculative-simple.cpp b/examples/speculative-simple/speculative-simple.cpp
index 08a1f2a88..e4a878b82 100644
--- a/examples/speculative-simple/speculative-simple.cpp
+++ b/examples/speculative-simple/speculative-simple.cpp
@@ -206,7 +206,7 @@ int main(int argc, char ** argv) {
// reset the draft context to the checkpoint before verification
if (ctx_dft) {
if (use_ckpt_dft) {
- ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ GGML_ASSERT(ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
}
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, ckpt.pos_max + 1, -1);
@@ -269,13 +269,13 @@ int main(int argc, char ** argv) {
draft = std::move(ids);
{
- ckpt.load_tgt(ctx_tgt, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ GGML_ASSERT(ckpt.load_tgt(ctx_tgt, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
llama_memory_seq_rm(llama_get_memory(ctx_tgt), seq_id, ckpt.pos_max + 1, -1);
}
if (ctx_dft) {
- ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ GGML_ASSERT(ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, ckpt.pos_max + 1, -1);
}
diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp
index e3270747b..1e3021d4a 100644
--- a/tools/server/server-context.cpp
+++ b/tools/server/server-context.cpp
@@ -2583,6 +2583,136 @@ private:
cur.pos_max, cur.n_tokens, (float) cur.size() / 1024 / 1024);
}
+ // checkpoints are appended to the slot save file, after the llama state payload
+ // they cannot be recreated from the final state alone (a recurrent state cannot be rewound)
+ static constexpr uint32_t SLOT_CKPT_MAGIC = 0x504b4353; // "SCKP"
+ static constexpr uint32_t SLOT_CKPT_VERSION = 1;
+
+ static bool ckpt_read(std::ifstream & ifs, void * dst, size_t size, size_t & n_read) {
+ if (!ifs.read((char *) dst, size)) {
+ return false;
+ }
+ n_read += size;
+ return true;
+ }
+
+ static bool ckpt_read_buf(std::ifstream & ifs, std::vector<uint8_t> & buf, size_t n_avail, size_t & n_read) {
+ uint64_t n = 0;
+ // check the size against the bytes left in the file before allocating, the size field may be corrupted
+ if (!ckpt_read(ifs, &n, sizeof(n), n_read) || n > n_avail - n_read) {
+ return false;
+ }
+ buf.resize(n);
+ return n == 0 || ckpt_read(ifs, buf.data(), n, n_read);
+ }
+
+ static void ckpt_write(std::ofstream & ofs, const void * src, size_t size, size_t & n_written) {
+ ofs.write((const char *) src, size);
+ n_written += size;
+ }
+
+ static void ckpt_write_buf(std::ofstream & ofs, const std::vector<uint8_t> & buf, size_t & n_written) {
+ const uint64_t n = buf.size();
+ ckpt_write(ofs, &n, sizeof(n), n_written);
+ if (n > 0) {
+ ckpt_write(ofs, buf.data(), n, n_written);
+ }
+ }
+
+ // returns false if the appendix could not be written completely
+ bool save_slot_checkpoints(const std::string & filepath, const server_slot & slot, size_t & n_written) const {
+ n_written = 0;
+ if (slot.prompt.checkpoints.empty()) {
+ return true;
+ }
+ std::ofstream ofs(std::filesystem::u8path(filepath), std::ios::binary | std::ios::app);
+ if (!ofs) {
+ SLT_WRN(slot, "failed to append context checkpoints to '%s'\n", filepath.c_str());
+ return false;
+ }
+ const uint32_t magic = SLOT_CKPT_MAGIC;
+ const uint32_t version = SLOT_CKPT_VERSION;
+ const uint32_t count = (uint32_t) slot.prompt.checkpoints.size();
+ ckpt_write(ofs, &magic, sizeof(magic), n_written);
+ ckpt_write(ofs, &version, sizeof(version), n_written);
+ ckpt_write(ofs, &count, sizeof(count), n_written);
+ for (const auto & cur : slot.prompt.checkpoints) {
+ ckpt_write(ofs, &cur.n_tokens, sizeof(cur.n_tokens), n_written);
+ ckpt_write(ofs, &cur.pos_min, sizeof(cur.pos_min), n_written);
+ ckpt_write(ofs, &cur.pos_max, sizeof(cur.pos_max), n_written);
+ ckpt_write_buf(ofs, cur.data_tgt, n_written);
+ ckpt_write_buf(ofs, cur.data_dft, n_written);
+ ckpt_write_buf(ofs, cur.data_spec, n_written);
+ }
+ ofs.flush();
+ if (!ofs) {
+ SLT_WRN(slot, "failed to append context checkpoints to '%s' - the appendix is incomplete\n", filepath.c_str());
+ return false;
+ }
+ SLT_INF(slot, "appended %u context checkpoint(s) (%.3f MiB) to '%s'\n",
+ count, (float) n_written / 1024 / 1024, filepath.c_str());
+ return true;
+ }
+
+ // returns the number of bytes consumed, 0 if there is no usable appendix
+ size_t load_slot_checkpoints(const std::string & filepath, size_t offset, server_slot & slot) const {
+ std::ifstream ifs(std::filesystem::u8path(filepath), std::ios::binary | std::ios::ate);
+ const size_t file_size = ifs ? (size_t) ifs.tellg() : 0;
+ if (!ifs || file_size < offset || !ifs.seekg(offset)) {
+ return 0;
+ }
+ const size_t n_avail = file_size - offset; // bytes after the llama state payload
+ size_t n_read = 0;
+ uint32_t magic = 0;
+ uint32_t version = 0;
+ uint32_t count = 0;
+ if (!ckpt_read(ifs, &magic, sizeof(magic), n_read) || magic != SLOT_CKPT_MAGIC) {
+ return 0;
+ }
+ if (!ckpt_read(ifs, &version, sizeof(version), n_read) || version != SLOT_CKPT_VERSION ||
+ !ckpt_read(ifs, &count, sizeof(count), n_read)) {
+ SLT_WRN(slot, "invalid context checkpoint appendix in '%s' - ignored\n", filepath.c_str());
+ return 0;
+ }
+ std::list<common_prompt_checkpoint> checkpoints;
+ for (uint32_t i = 0; i < count; ++i) {
+ common_prompt_checkpoint cur;
+ cur.id_task = -1; // not created by a task - marks a checkpoint restored from a slot file
+ if (!ckpt_read(ifs, &cur.n_tokens, sizeof(cur.n_tokens), n_read) ||
+ !ckpt_read(ifs, &cur.pos_min, sizeof(cur.pos_min), n_read) ||
+ !ckpt_read(ifs, &cur.pos_max, sizeof(cur.pos_max), n_read) ||
+ !ckpt_read_buf(ifs, cur.data_tgt, n_avail, n_read) ||
+ !ckpt_read_buf(ifs, cur.data_dft, n_avail, n_read) ||
+ !ckpt_read_buf(ifs, cur.data_spec, n_avail, n_read)) {
+ SLT_WRN(slot, "truncated context checkpoint appendix in '%s' - ignored\n", filepath.c_str());
+ return 0;
+ }
+ // a saved checkpoint always holds a target state - an empty blob would roll back without restoring anything
+ if (cur.data_tgt.empty()) {
+ SLT_WRN(slot, "invalid context checkpoint appendix in '%s' - ignored\n", filepath.c_str());
+ return 0;
+ }
+ checkpoints.push_back(std::move(cur));
+ if (checkpoints.size() > (size_t) params_base.n_ctx_checkpoints) {
+ checkpoints.pop_front();
+ }
+ }
+ // the slot file does not check the draft context - test-load one draft checkpoint, drop the draft data if it does not fit
+ if (ctx_dft != nullptr && !checkpoints.empty() && !checkpoints.back().data_dft.empty()) {
+ const bool ok = checkpoints.back().load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ llama_memory_seq_rm(llama_get_memory(ctx_dft), slot.id, -1, -1);
+ if (!ok) {
+ SLT_WRN(slot, "draft context checkpoint data in '%s' does not match the draft context - dropped\n", filepath.c_str());
+ for (auto & cur : checkpoints) {
+ cur.clear_dft();
+ }
+ }
+ }
+ slot.prompt.checkpoints = std::move(checkpoints);
+ SLT_INF(slot, "restored %zu context checkpoint(s) from '%s'\n", slot.prompt.checkpoints.size(), filepath.c_str());
+ return n_read;
+ }
+
// returns false to decline the task, it is offered again after the decode is done
bool process_single_task(server_task && task, bool is_yielding) {
// while yielding, an encode / decode is running and only reading the server state is safe
@@ -2792,6 +2922,12 @@ private:
break;
}
+ size_t nwrite_ckpt = 0;
+ if (!save_slot_checkpoints(filepath, *slot, nwrite_ckpt)) {
+ send_error(task, "Unable to save slot: incomplete context checkpoints", ERROR_TYPE_SERVER);
+ break;
+ }
+
const int64_t t_end = ggml_time_us();
const double t_save_ms = (t_end - t_start) / 1000.0;
@@ -2801,7 +2937,7 @@ private:
res->filename = filename;
res->is_save = true;
res->n_tokens = slot->prompt.tokens.size();
- res->n_bytes = nwrite;
+ res->n_bytes = nwrite + nwrite_ckpt;
res->t_ms = t_save_ms;
queue_results.send(std::move(res));
} break;
@@ -2857,6 +2993,9 @@ private:
break;
}
+ // nread is the end offset of the llama state payload within the file
+ const size_t nread_ckpt = load_slot_checkpoints(filepath, nread, *slot);
+
const int64_t t_end = ggml_time_us();
const double t_restore_ms = (t_end - t_start) / 1000.0;
@@ -2866,7 +3005,7 @@ private:
res->filename = filename;
res->is_save = false;
res->n_tokens = slot->prompt.tokens.size();
- res->n_bytes = nread;
+ res->n_bytes = nread + nread_ckpt;
res->t_ms = t_restore_ms;
queue_results.send(std::move(res));
} break;
@@ -3278,7 +3417,7 @@ private:
if (ctx_dft) {
if (use_ckpt_dft) {
- ckpt.load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ GGML_ASSERT(ckpt.load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
}
if (!llama_memory_seq_rm(llama_get_memory(ctx_dft), slot.id, ckpt.pos_max + 1, -1)) {
@@ -3604,8 +3743,18 @@ private:
if (!do_reset) {
// restore the context checkpoint
- it->load_tgt(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
- it->load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ if (!it->load_tgt(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY) ||
+ !it->load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY)) {
+ if (it->id_task != -1) {
+ GGML_ABORT("failed to restore context checkpoint\n");
+ }
+ // restored from a slot file, not guaranteed to load - fall back to full prompt re-processing
+ SLT_WRN(slot, "%s", "failed to load context checkpoint restored from a slot file\n");
+ do_reset = true;
+ }
+ }
+
+ if (!do_reset) {
// restore the draft's speculative state
common_speculative_set_state(spec.get(), slot.id, it->data_spec);
@@ -4300,10 +4449,10 @@ private:
SLT_DBG(slot, "restoring speculative checkpoint (pos_min = %d, pos_max = %d, size = %zu)\n", ckpt.pos_min, ckpt.pos_max, ckpt.size());
- ckpt.load_tgt(slot.ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ GGML_ASSERT(ckpt.load_tgt(slot.ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
if (slot.ctx_dft) {
- ckpt.load_dft(slot.ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
+ GGML_ASSERT(ckpt.load_dft(slot.ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
}
slot.mem.seq_rm(slot.id, ckpt.pos_max + 1, -1);
diff --git a/tools/server/tests/unit/test_slot_save.py b/tools/server/tests/unit/test_slot_save.py
index 5eca46cb2..33bcede93 100644
--- a/tools/server/tests/unit/test_slot_save.py
+++ b/tools/server/tests/unit/test_slot_save.py
@@ -37,6 +37,8 @@ def test_slot_save_restore():
})
assert res.status_code == 200
assert res.body["n_saved"] == 84
+ slot_file = os.path.join(server.slot_save_path, "slot1.bin")
+ assert res.body["n_written"] == os.path.getsize(slot_file)
# Since we have cache, this should only process the last tokens
res = server.make_request("POST", "/completion", data={
@@ -54,6 +56,7 @@ def test_slot_save_restore():
})
assert res.status_code == 200
assert res.body["n_restored"] == 84
+ assert res.body["n_read"] == os.path.getsize(slot_file)
# Since we have cache, slot 0 should only process the last tokens
res = server.make_request("POST", "/completion", data={
@@ -546,3 +549,243 @@ def test_slot_restore_media_file_without_mmproj(mmproj_server):
assert res.status_code == 200
assert res.body["timings"]["cache_n"] == 0
assert res.body["content"] == content
+
+
+@pytest.fixture
+def swa_server():
+ swa = ServerPreset.tinygemma3()
+ swa.slot_save_path = "./tmp"
+ swa.temperature = 0.0
+ swa.cache_ram = 0
+ # Keep the first prompt checkpoint before the divergence point.
+ swa.n_ubatch = 32
+ return swa
+
+
+# the non-ASCII name checks that the appendix lands in the same file as the llama state on Windows
+@pytest.mark.parametrize("filename", ["ckpt_slot1.bin", "ckpt_slot1_é.bin"])
+def test_slot_restore_preserves_context_checkpoints(swa_server, filename):
+ server = swa_server
+ server.start()
+
+ base = "The quick brown fox jumps over the lazy dog. " * 20
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "The first ending of this story is a happy one.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+ n_full = res.body["timings"]["prompt_n"]
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "But the second ending was different and sad.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+ n_live = res.body["timings"]["prompt_n"]
+ assert n_live < n_full
+
+ res = server.make_request("POST", "/slots/1?action=erase")
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "The first ending of this story is a happy one.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/slots/1?action=save", data={
+ "filename": filename,
+ })
+ assert res.status_code == 200
+ assert res.body["n_saved"] > 0
+ ckpt_file = os.path.join(server.slot_save_path, filename)
+ assert res.body["n_written"] == os.path.getsize(ckpt_file)
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": "Unrelated text with no common prefix occupies the slot now.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/slots/1?action=restore", data={
+ "filename": filename,
+ })
+ assert res.status_code == 200
+ assert res.body["n_read"] == os.path.getsize(ckpt_file)
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "But the second ending was different and sad.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+ assert res.body["timings"]["prompt_n"] == n_live
+
+
+# checkpoint appendix: magic(4) version(4) count(4), then per checkpoint
+# n_tokens(8) pos_min(4) pos_max(4) and three blobs (target, draft, speculative), each size(8) + data
+def parse_ckpt_appendix(data):
+ off = data.find(struct.pack("<II", 0x504b4353, 1))
+ assert off > 0
+ count = struct.unpack_from("<I", data, off + 8)[0]
+ ckpts = []
+ pos = off + 12
+ for _ in range(count):
+ start = pos
+ pos += 16
+ blobs = []
+ for _ in range(3):
+ n = struct.unpack_from("<Q", data, pos)[0]
+ blobs.append(pos + 8)
+ pos += 8 + n
+ ckpts.append((start, pos, blobs[0]))
+ assert pos == len(data)
+ return off, ckpts
+
+
+# a damaged appendix must be ignored, or its checkpoints dropped when they fail to load, without aborting the server
+@pytest.mark.parametrize("damage", ["oversized_blob", "empty_target", "corrupt_state", "many_checkpoints"])
+def test_slot_restore_damaged_checkpoint_appendix(swa_server, damage):
+ server = swa_server
+ server.start()
+
+ base = "The quick brown fox jumps over the lazy dog. " * 20
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "The first ending of this story is a happy one.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "But the second ending was different and sad.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+ n_live = res.body["timings"]["prompt_n"]
+
+ res = server.make_request("POST", "/slots/1?action=erase")
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "The first ending of this story is a happy one.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/slots/1?action=save", data={
+ "filename": "ckpt_damaged.bin",
+ })
+ assert res.status_code == 200
+
+ path = os.path.join(server.slot_save_path, "ckpt_damaged.bin")
+ with open(path, "rb") as f:
+ data = bytearray(f.read())
+ off, ckpts = parse_ckpt_appendix(data)
+
+ if damage == "oversized_blob":
+ # the first target blob declares a size that cannot be allocated, it must be rejected before allocating
+ data = data[:ckpts[0][0] + 16] + struct.pack("<Q", 1 << 62)
+ elif damage == "empty_target":
+ # the target blobs are removed and their size set to 0, a valid save never writes an empty target state
+ for start, end, tgt in reversed(ckpts):
+ size = struct.unpack_from("<Q", data, tgt - 8)[0]
+ data = data[:tgt - 8] + struct.pack("<Q", 0) + data[tgt + size:]
+ elif damage == "corrupt_state":
+ # the sizes are intact, but the target states do not load
+ for _, _, tgt in ckpts:
+ struct.pack_into("<I", data, tgt, 0xdeadbeef)
+ else:
+ # more than 1024 entries: one-byte fillers that never match go first, the real checkpoints stay last
+ filler = struct.pack("<qiiQBQQ", 0, 0, 1 << 30, 1, 0, 0, 0)
+ data = data[:off + 12] + filler * (1025 - len(ckpts)) + data[off + 12:]
+ struct.pack_into("<I", data, off + 8, 1025)
+
+ with open(path, "wb") as f:
+ f.write(data)
+
+ res = server.make_request("POST", "/slots/1?action=restore", data={
+ "filename": "ckpt_damaged.bin",
+ })
+ assert res.status_code == 200
+ if damage in ("oversized_blob", "empty_target"):
+ assert res.body["n_read"] == off
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "But the second ending was different and sad.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+ if damage == "many_checkpoints":
+ assert res.body["timings"]["prompt_n"] == n_live
+ else:
+ assert res.body["timings"]["prompt_n"] > n_live
+
+
+# the draft blobs of the checkpoint appendix are not covered by the main payload checks,
+# so restoring into a server with another draft KV cache type must not abort
+@pytest.mark.parametrize("ctkd_restore", ["f16", "q8_0"])
+def test_slot_restore_checkpoints_draft_kv_type_change(swa_server, ctkd_restore):
+ server = swa_server
+ server.model_draft_hf_repo = "ggml-org/tinygemma3-GGUF:Q8_0" # same file as the target, already in the HF cache
+ server.spec_type = "draft-simple"
+ server.ctkd = "f16"
+ server.start()
+
+ base = "The quick brown fox jumps over the lazy dog. " * 20
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "The first ending of this story is a happy one.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "But the second ending was different and sad.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+ n_live = res.body["timings"]["prompt_n"]
+
+ res = server.make_request("POST", "/slots/1?action=erase")
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "The first ending of this story is a happy one.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/slots/1?action=save", data={
+ "filename": "ckpt_draft_slot1.bin",
+ })
+ assert res.status_code == 200
+
+ server.stop()
+ server.ctkd = ctkd_restore
+ server.start()
+
+ res = server.make_request("POST", "/slots/1?action=restore", data={
+ "filename": "ckpt_draft_slot1.bin",
+ })
+ assert res.status_code == 200
+
+ res = server.make_request("POST", "/completion", data={
+ "prompt": base + "But the second ending was different and sad.",
+ "id_slot": 1,
+ "cache_prompt": True,
+ })
+ assert res.status_code == 200
+ assert res.body["timings"]["prompt_n"] == n_live
diff --git a/tools/server/tests/utils.py b/tools/server/tests/utils.py
index 90c4ebac8..76e7b2bf4 100644
--- a/tools/server/tests/utils.py
+++ b/tools/server/tests/utils.py
@@ -65,6 +65,7 @@ class ServerProcess:
model_url: str | None = None
model_file: str | None = None
model_draft: str | None = None
+ model_draft_hf_repo: str | None = None
n_threads: int | None = None
n_gpu_layer: int | None = None
n_batch: int | None = None
@@ -80,6 +81,7 @@ class ServerProcess:
n_slots: int | None = None
ctk: str | None = None
ctv: str | None = None
+ ctkd: str | None = None
fa: str | None = None
server_continuous_batching: bool | None = False
server_embeddings: bool | None = False
@@ -171,6 +173,8 @@ class ServerProcess:
server_args.extend(["--model-url", self.model_url])
if self.model_draft:
server_args.extend(["--model-draft", self.model_draft])
+ if self.model_draft_hf_repo:
+ server_args.extend(["--hf-repo-draft", self.model_draft_hf_repo])
if self.model_hf_repo:
server_args.extend(["--hf-repo", self.model_hf_repo])
if self.model_hf_file:
@@ -221,6 +225,8 @@ class ServerProcess:
server_args.extend(["-ctk", self.ctk])
if self.ctv:
server_args.extend(["-ctv", self.ctv])
+ if self.ctkd:
+ server_args.extend(["-ctkd", self.ctkd])
if self.fa is not None:
server_args.extend(["-fa", self.fa])
if self.n_predict: