diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 7d636c7923b8..af28bacdea62 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -1808,8 +1808,8 @@ struct server_context_impl { f_sim_best, slot_prompt_similarity, f_keep); } - // if we are about to lose a large portion of the existing context - save it in the prompt cache - if (f_keep < 0.5f) { + // A replaced branch can share most of a tool schema and still need its own cached suffix. + if (ret->prompt.tokens.get_common_prefix(task.tokens) < ret->prompt.tokens.size()) { update_cache = true; } } diff --git a/tools/server/tests/unit/test_slot_save.py b/tools/server/tests/unit/test_slot_save.py index 5eca46cb292d..6e07bdc43c5c 100644 --- a/tools/server/tests/unit/test_slot_save.py +++ b/tools/server/tests/unit/test_slot_save.py @@ -158,6 +158,39 @@ def test_slot_erase(): assert res.body["timings"]["prompt_n"] == 21 # all tokens are processed +def test_ram_cache_interleaved_shared_prefix(): + server.n_slots = 1 + server.cache_ram = 16 + server.n_predict = 8 + server.start() + + prefix = [1] + [10] * 192 + prompts = [prefix + [20] * 48, prefix + [30] * 48] + + def complete(prompt): + res = server.make_request("POST", "/completion", data={ + "prompt": prompt, + "cache_prompt": True, + "n_predict": 8, + "ignore_eos": True, + "return_tokens": True, + "temperature": 0.0, + }) + assert res.status_code == 200 + return res.body + + first = [] + for prompt in prompts: + complete(prompt) + # Compare the same one-token replay shape before and after displacement. + first.append(complete(prompt)) + for prompt, expected in zip(prompts, first): + restored = complete(prompt) + assert restored["timings"]["cache_n"] == len(prompt) - 1 + assert restored["timings"]["prompt_n"] == 1 + assert restored["tokens"] == expected["tokens"] + + # # Multimodal server (mmproj loaded) slot save/restore. #