diff --git a/common/arg.cpp b/common/arg.cpp index a036b44e..2487497d 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -1724,6 +1724,42 @@ common_params_context common_params_parser_init(common_params & params, llama_ex params.cache_disk_mib = value; } ).set_env("LLAMA_ARG_CACHE_DISK_SIZE").set_examples({LLAMA_EXAMPLE_SERVER})); + add_opt(common_arg( + {"--cache-ttl-ephemeral"}, "N", + string_format("seconds a prompt cache entry of class ephemeral lives after its last use (default: %d)", params.cache_ttl_ephemeral), + [](common_params & params, int value) { + params.cache_ttl_ephemeral = std::max(value, 0); + } + ).set_env("LLAMA_ARG_CACHE_TTL_EPHEMERAL").set_examples({LLAMA_EXAMPLE_SERVER})); + add_opt(common_arg( + {"--cache-ttl-default"}, "N", + string_format("seconds a prompt cache entry of class default lives after its last use (default: %d)", params.cache_ttl_default), + [](common_params & params, int value) { + params.cache_ttl_default = std::max(value, 0); + } + ).set_env("LLAMA_ARG_CACHE_TTL_DEFAULT").set_examples({LLAMA_EXAMPLE_SERVER})); + add_opt(common_arg( + {"--cache-ttl-keep"}, "N", + string_format("seconds a prompt cache entry of class keep lives after its last use (default: %d)", params.cache_ttl_keep), + [](common_params & params, int value) { + params.cache_ttl_keep = std::max(value, 0); + } + ).set_env("LLAMA_ARG_CACHE_TTL_KEEP").set_examples({LLAMA_EXAMPLE_SERVER})); + add_opt(common_arg( + {"--cache-classifier"}, "URL", + "POST new prompt cache entries without an explicit class to URL as {excerpt, n_tokens, n_turns};\n" + "it answers {class, ttl?} (default: off)", + [](common_params & params, const std::string & value) { + params.cache_classifier = value; + } + ).set_env("LLAMA_ARG_CACHE_CLASSIFIER").set_examples({LLAMA_EXAMPLE_SERVER})); + add_opt(common_arg( + {"--cache-classifier-timeout"}, "N", + string_format("seconds to wait for the cache classifier (default: %d)", params.cache_classifier_timeout), + [](common_params & params, int value) { + params.cache_classifier_timeout = std::max(value, 1); + } + ).set_env("LLAMA_ARG_CACHE_CLASSIFIER_TIMEOUT").set_examples({LLAMA_EXAMPLE_SERVER})); add_opt(common_arg( {"-kvu", "--kv-unified"}, {"-no-kvu", "--no-kv-unified"}, diff --git a/common/common.h b/common/common.h index e70f2a99..76028932 100644 --- a/common/common.h +++ b/common/common.h @@ -621,6 +621,11 @@ struct common_params { int32_t cache_ram_mib = 8192; // -1 = no limit, 0 - disable, 1 = 1 MiB, etc. std::string cache_disk_path; // directory of the prompt cache's disk tier, empty = off int32_t cache_disk_mib = 51200; // disk tier size limit in MiB, 0 = no limit + int32_t cache_ttl_ephemeral = 3600; // seconds an unused prompt cache entry of each class is kept + int32_t cache_ttl_default = 3*86400; + int32_t cache_ttl_keep = 30*86400; + std::string cache_classifier; // URL that classifies new prompt cache entries, empty = off + int32_t cache_classifier_timeout = 20; // seconds std::string hostname = "127.0.0.1"; std::string public_path = ""; // NOLINT diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 8755886b..bed091d9 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -251,8 +251,11 @@ struct server_slot { server_prompt prompt; + // the request that produced the prompt asked for "cache": "skip" + bool cache_skip = false; + bool prompt_save(server_prompt_cache & prompt_cache) const { - if (prompt.tokens.size() == 0) { + if (prompt.tokens.size() == 0 || cache_skip) { return false; } @@ -277,8 +280,8 @@ struct server_slot { return true; } - bool prompt_load(server_prompt_cache & prompt_cache, const server_tokens & tokens) { - bool res = prompt_cache.load(prompt, tokens, ctx_tgt, ctx_dft, id); + bool prompt_load(server_prompt_cache & prompt_cache, const server_tokens & tokens, bool keep_entry) { + bool res = prompt_cache.load(prompt, tokens, ctx_tgt, ctx_dft, id, keep_entry); if (!res) { SLT_WRN(*this, "%s", "failed to load prompt from cache\n"); } @@ -892,6 +895,7 @@ private: void destroy() { // the slots and the RAM cache go to the disk tier while the contexts can still read them if (prompt_cache) { + prompt_cache->maintain(slot_metas()); for (auto & slot : slots) { slot.prompt_save(*prompt_cache); } @@ -913,6 +917,14 @@ private: mctx = nullptr; } + std::vector slot_metas() { + std::vector res; + for (auto & slot : slots) { + res.push_back(&slot.prompt.meta); + } + return res; + } + void handle_sleeping_state(bool new_state) { GGML_ASSERT(sleeping != new_state); if (new_state) { @@ -1318,6 +1330,30 @@ private: prompt_cache = std::make_unique(params_base.cache_ram_mib, n_ctx); + { + cache_policy policy; + policy.ttl_ephemeral = params_base.cache_ttl_ephemeral; + policy.ttl_default = params_base.cache_ttl_default; + policy.ttl_keep = params_base.cache_ttl_keep; + policy.classifier_url = params_base.cache_classifier; + policy.classifier_timeout = params_base.cache_classifier_timeout; + const llama_vocab * vocab = llama_model_get_vocab(model_tgt); + prompt_cache->policy_init(policy, + // media positions are LLAMA_TOKEN_NULL and have no text + [vocab](const llama_tokens & tokens) { + llama_tokens text; + std::copy_if(tokens.begin(), tokens.end(), std::back_inserter(text), + [](llama_token t) { return t != LLAMA_TOKEN_NULL; }); + return common_detokenize(vocab, text, true); + }, + common_tokenize(vocab, "<|im_start|>user", false, true), + [this]() { + server_task task(SERVER_TASK_TYPE_CACHE_MAINT); + task.id = queue_tasks.get_new_id(); + queue_tasks.post(std::move(task)); + }); + } + if (!params_base.cache_disk_path.empty()) { // a file is only valid for the same weights (path, size, mtime) and cache types std::error_code ec; @@ -1330,6 +1366,7 @@ private: ggml_type_name(params_base.speculative.draft.cache_type_k), ggml_type_name(params_base.speculative.draft.cache_type_v)); prompt_cache->disk_init(params_base.cache_disk_path, params_base.cache_disk_mib, key); } + prompt_cache->policy_start(); } else { SRV_TRC("%s", "prompt cache is disabled - use `--cache-ram N` to enable it\n"); } @@ -1601,14 +1638,25 @@ private: // cache prompts only for completion tasks update_cache = update_cache && task.type == SERVER_TASK_TYPE_COMPLETION; + const bool skip = task.params.cache_explicit && task.params.cache_cls == CACHE_CLASS_SKIP; + + // a skip request must not take the slot's unsaved conversation with it + if (!update_cache && skip && prompt_cache && task.type == SERVER_TASK_TYPE_COMPLETION && !ret->cache_skip) { + if (ret->prompt_save(*prompt_cache)) { + prompt_cache->update(); + } + } + if (update_cache) { SRV_TRC("%s", "updating prompt cache\n"); const int64_t t_start = ggml_time_us(); + prompt_cache->maintain(slot_metas()); + ret->prompt_save(*prompt_cache); - if (!ret->prompt_load(*prompt_cache, task.tokens)) { + if (!ret->prompt_load(*prompt_cache, task.tokens, skip)) { ret->prompt_clear(); } @@ -1684,6 +1732,20 @@ private: slot.lora = params_base.lora_adapters; } + // a continuation covers the previous request's prompt (all but a few template tokens at its end); + // without that record (v1 files), the f_keep >= 0.25 rule of the prompt cache + if (prompt_cache && task.type == SERVER_TASK_TYPE_COMPLETION) { + const auto & meta = slot.prompt.meta; + const int64_t lcp = slot.prompt.tokens.get_common_prefix(task.tokens); + const bool continuation = !slot.prompt.tokens.empty() && lcp > 0 && (meta.n_input > 0 + ? lcp + 32 >= meta.n_input && 2*lcp >= meta.n_input + : 4*lcp >= (int64_t) slot.prompt.tokens.size()); + prompt_cache->apply_request(slot.prompt.meta, continuation, task.params, task.tokens.size()); + slot.cache_skip = task.params.cache_explicit && task.params.cache_cls == CACHE_CLASS_SKIP; + SLT_INF(slot, "cache: %s, class %s (%s), hits %llu%s\n", continuation ? "continues the slot prompt" : "new prompt", + cache_class_name(meta.cls), cache_source_name(meta.source), (unsigned long long) meta.hits, slot.cache_skip ? ", skip" : ""); + } + // if using alora, make sure it's only a single one requested and active size_t alora_invocation_start = task.tokens.size(); if (lora_all_alora(slot.lora)) { @@ -2632,10 +2694,63 @@ private: res->n_erased = n_erased; queue_results.send(std::move(res)); } break; + case SERVER_TASK_TYPE_CACHE_MAINT: + { + if (prompt_cache) { + prompt_cache->maintain(slot_metas()); + } + } break; case SERVER_TASK_TYPE_CACHE_GET: case SERVER_TASK_TYPE_CACHE_CLEAR: + case SERVER_TASK_TYPE_CACHE_ENTRY: { + if (prompt_cache) { + prompt_cache->maintain(slot_metas()); + } json cleared = json::object(); + json entry = json::object(); + if (task.type == SERVER_TASK_TYPE_CACHE_ENTRY) { + if (!prompt_cache) { + send_error(task, "the prompt cache is off", ERROR_TYPE_NOT_SUPPORTED); + break; + } + uint64_t id = 0; + try { + id = std::stoull(task.cache_entry_id, nullptr, 16); + } catch (const std::exception &) { + send_error(task, "id must be an entry id from GET /cache", ERROR_TYPE_INVALID_REQUEST); + break; + } + bool found = false; + if (task.cache_entry_action == "delete") { + found = prompt_cache->delete_entry(id); + } else if (task.cache_entry_action.empty()) { + cache_class cls; + if (!cache_class_from_name(task.cache_entry_class, cls) || cls == CACHE_CLASS_SKIP) { + send_error(task, "class must be ephemeral, default, keep or pin", ERROR_TYPE_INVALID_REQUEST); + break; + } + found = prompt_cache->set_entry(id, cls, task.cache_entry_ttl); + // the entry may have just moved into a slot + for (auto & slot : slots) { + if (slot.prompt.meta.id == id && !slot.prompt.tokens.empty()) { + slot.prompt.meta.cls = cls; + slot.prompt.meta.source = CACHE_SOURCE_EXPLICIT; + slot.prompt.meta.ttl = task.cache_entry_ttl; + found = true; + } + } + } else { + send_error(task, "action must be delete, or omitted with class", ERROR_TYPE_INVALID_REQUEST); + break; + } + if (!found) { + send_error(task, "no cache entry with id " + task.cache_entry_id, ERROR_TYPE_NOT_FOUND); + break; + } + entry = { { "id", task.cache_entry_id }, { "action", task.cache_entry_action.empty() ? task.cache_entry_class : task.cache_entry_action } }; + SRV_INF("cache entry %s: %s\n", task.cache_entry_id.c_str(), entry.dump().c_str()); + } if (task.type == SERVER_TASK_TYPE_CACHE_CLEAR) { const std::string & scope = task.cache_scope; const bool all = scope == "all"; @@ -2650,7 +2765,8 @@ private: for (auto & slot : slots) { if (slot.is_processing()) { n_busy++; - } else if (!slot.prompt.tokens.empty()) { + } else if (!slot.prompt.tokens.empty() && + (task.cache_include_pinned || slot.prompt.meta.cls != CACHE_CLASS_PIN)) { slot.prompt_clear(); n_cleared++; } @@ -2659,11 +2775,10 @@ private: cleared["slots_busy"] = n_busy; } if (prompt_cache && (all || scope == "ram")) { - cleared["ram"] = prompt_cache->states.size(); - prompt_cache->states.clear(); + cleared["ram"] = prompt_cache->clear_ram(task.cache_include_pinned); } if (prompt_cache && (all || scope == "disk")) { - cleared["disk"] = prompt_cache->clear_disk(); + cleared["disk"] = prompt_cache->clear_disk(task.cache_include_pinned); } SRV_INF("cache cleared (%s): %s\n", scope.c_str(), cleared.dump().c_str()); } @@ -2674,6 +2789,10 @@ private: { "id", slot.id }, { "n_tokens", slot.prompt.tokens.size() }, { "processing", slot.is_processing() }, + { "class", cache_class_name(slot.cache_skip ? CACHE_CLASS_SKIP : slot.prompt.meta.cls) }, + { "source", cache_source_name(slot.prompt.meta.source) }, + { "hits", slot.prompt.meta.hits }, + { "excerpt", slot.prompt.meta.excerpt }, }); } auto res = std::make_unique(); @@ -2683,6 +2802,9 @@ private: if (task.type == SERVER_TASK_TYPE_CACHE_CLEAR) { res->data["cleared"] = cleared; } + if (task.type == SERVER_TASK_TYPE_CACHE_ENTRY) { + res->data["entry"] = entry; + } queue_results.send(std::move(res)); } break; case SERVER_TASK_TYPE_GET_LORA: @@ -4771,6 +4893,11 @@ void server_routes::init_routes() { return handle_cache(req, SERVER_TASK_TYPE_CACHE_CLEAR, scope.empty() ? "all" : scope); }; + // POST /cache/entry?id=ID&class=CLASS[&ttl=SECONDS] or ?id=ID&action=delete + this->post_cache_entry = [this](const server_http_req & req) { + return handle_cache(req, SERVER_TASK_TYPE_CACHE_ENTRY, ""); + }; + this->post_slots = [this](const server_http_req & req) { auto res = create_response(); if (params.slot_save_path.empty()) { @@ -5370,6 +5497,20 @@ std::unique_ptr server_routes::handle_cache(const server_h server_task task(type); task.id = rd.get_new_id(); task.cache_scope = scope; + const std::string pinned = req.get_param("include_pinned"); + task.cache_include_pinned = pinned == "1" || pinned == "true"; + task.cache_entry_id = req.get_param("id"); + task.cache_entry_class = req.get_param("class"); + task.cache_entry_action = req.get_param("action"); + const std::string ttl = req.get_param("ttl"); + if (!ttl.empty()) { + try { + task.cache_entry_ttl = std::max(0, std::stoll(ttl)); + } catch (const std::exception &) { + res->error(format_error_response("ttl must be a number of seconds", ERROR_TYPE_INVALID_REQUEST)); + return res; + } + } rd.post_task(std::move(task), true); // high-priority task } diff --git a/tools/server/server-context.h b/tools/server/server-context.h index adb1fa81..fa342d3a 100644 --- a/tools/server/server-context.h +++ b/tools/server/server-context.h @@ -134,6 +134,7 @@ struct server_routes { server_http_context::handler_t post_slots; server_http_context::handler_t get_cache; server_http_context::handler_t post_cache_clear; + server_http_context::handler_t post_cache_entry; server_http_context::handler_t get_props; server_http_context::handler_t post_props; server_http_context::handler_t post_infill; diff --git a/tools/server/server-queue.cpp b/tools/server/server-queue.cpp index 78169e9a..b6dd8fc0 100644 --- a/tools/server/server-queue.cpp +++ b/tools/server/server-queue.cpp @@ -22,7 +22,7 @@ // static bool task_resets_idle_timer(server_task_type type) { - return type != SERVER_TASK_TYPE_METRICS; + return type != SERVER_TASK_TYPE_METRICS && type != SERVER_TASK_TYPE_CACHE_MAINT; } int server_queue::post(server_task && task, bool front) { diff --git a/tools/server/server-schema.cpp b/tools/server/server-schema.cpp index 64b92512..8254e4a8 100644 --- a/tools/server/server-schema.cpp +++ b/tools/server/server-schema.cpp @@ -31,6 +31,20 @@ std::vector> make_llama_cmpl_schema(const common_params & add((new field_bool("cache_prompt", params.cache_prompt)) ->set_desc("Re-use KV cache from a previous request if possible. This way the common prefix does not have to be re-processed, only the suffix that differs between the requests")); + add((new field_str("cache")) + ->set_desc("Retention of the prompt this request leaves in the prompt cache: skip, ephemeral, default, keep or pin") + ->set_handler([&](field_eval_context & ctx, const json & data) { + const std::string name = data.at("cache").get(); + if (!cache_class_from_name(name, ctx.params.cache_cls)) { + throw std::invalid_argument("must be skip, ephemeral, default, keep or pin"); + } + ctx.params.cache_explicit = true; + })); + + add((new field_num("cache_ttl", params.cache_ttl)) + ->set_hard_limits(0, std::numeric_limits::max()) + ->set_desc("Seconds the prompt stays in the prompt cache after its last use, instead of its class default")); + add((new field_bool("return_tokens", params.return_tokens)) ->set_desc("Return the raw generated token ids in the `tokens` field")); diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp index ea3b0c81..ad6565c7 100644 --- a/tools/server/server-task.cpp +++ b/tools/server/server-task.cpp @@ -14,6 +14,13 @@ #include #include #include +#include +#include +#include +#include +#include + +#include // // task_params @@ -1691,10 +1698,69 @@ json server_task_result_apply_lora::to_json() { // // server_prompt_cache // + +const char * cache_class_name(cache_class cls) { + switch (cls) { + case CACHE_CLASS_SKIP: return "skip"; + case CACHE_CLASS_EPHEMERAL: return "ephemeral"; + case CACHE_CLASS_DEFAULT: return "default"; + case CACHE_CLASS_KEEP: return "keep"; + case CACHE_CLASS_PIN: return "pin"; + } + return "default"; +} + +bool cache_class_from_name(const std::string & name, cache_class & out) { + for (cache_class c : { CACHE_CLASS_SKIP, CACHE_CLASS_EPHEMERAL, CACHE_CLASS_DEFAULT, CACHE_CLASS_KEEP, CACHE_CLASS_PIN }) { + if (name == cache_class_name(c)) { + out = c; + return true; + } + } + return false; +} + +const char * cache_source_name(cache_source src) { + switch (src) { + case CACHE_SOURCE_DEFAULT: return "default"; + case CACHE_SOURCE_CLASSIFIER: return "classifier"; + case CACHE_SOURCE_FAMILY: return "family"; + case CACHE_SOURCE_EXPLICIT: return "explicit"; + } + return "default"; +} + +int64_t cache_policy::ttl(cache_class cls) const { + switch (cls) { + case CACHE_CLASS_EPHEMERAL: return ttl_ephemeral; + case CACHE_CLASS_KEEP: return ttl_keep; + case CACHE_CLASS_PIN: return 0; + default: return ttl_default; + } +} + namespace { constexpr uint32_t DISK_MAGIC = 0x43564b42; // "BKVC" -constexpr uint32_t DISK_VERSION = 1; +constexpr uint32_t DISK_VERSION = 2; + +// v2: magic, version, key, META_SIZE bytes of metadata (rewritten in place), then the v1 body +constexpr size_t META_OFFSET = 16; +constexpr size_t META_SIZE = 512; +constexpr size_t META_FIXED = 74; + +constexpr size_t FAMILY_TOKENS = 256; +constexpr size_t FAMILY_MIN_N = 8; +constexpr size_t FAMILIES_MAX = 500; +constexpr size_t CLASSIFY_HEAD = 1500; +constexpr size_t CLASSIFY_TAIL = 500; +constexpr size_t CLASSIFY_QUEUE_MAX = 64; +constexpr size_t EXCERPT_CHARS = 80; +constexpr size_t INFO_ENTRIES = 200; + +int64_t unix_now() { + return (int64_t) std::time(nullptr); +} uint64_t fnv1a(const void * data, size_t n, uint64_t h = 1469598103934665603ull) { const auto * p = static_cast(data); @@ -1704,6 +1770,189 @@ uint64_t fnv1a(const void * data, size_t n, uint64_t h = 1469598103934665603ull) return h; } +uint64_t new_entry_id() { + static std::mt19937_64 rng(std::random_device{}() ^ (uint64_t) ggml_time_us()); + uint64_t id = 0; + while (id == 0) { + id = rng(); + } + return id; +} + +std::string hex_id(uint64_t id) { + return string_format("%016llx", (unsigned long long) id); +} + +// the longest valid UTF-8 prefix of s, invalid bytes replaced, cut at a character boundary +std::string clean_utf8(const std::string & s, size_t max_bytes) { + std::string out; + size_t i = 0; + while (i < s.size()) { + const auto c = (unsigned char) s[i]; + size_t n = c < 0x80 ? 1 : (c >> 5) == 0x6 ? 2 : (c >> 4) == 0xe ? 3 : (c >> 3) == 0x1e ? 4 : 0; + bool ok = n > 0 && i + n <= s.size(); + for (size_t k = 1; ok && k < n; ++k) { + ok = ((unsigned char) s[i + k] >> 6) == 0x2; + } + const size_t len = ok ? n : 1; + if (out.size() + len > max_bytes) { + break; + } + if (ok) { + out.append(s, i, n); + } else { + out.push_back('?'); + } + i += len; + } + return out; +} + +// the start of the first user message of a rendered chat prompt (special tokens as text), or of the prompt +std::string user_excerpt(const std::string & text) { + static const char * markers[] = { + "<|im_start|>user\n", "user\n", "<|start_header_id|>user<|end_header_id|>\n\n", "<|user|>\n", "[INST]", + }; + size_t begin = std::string::npos; + for (const char * m : markers) { + const size_t p = text.find(m); + if (p != std::string::npos) { + begin = p + strlen(m); + break; + } + } + std::string res = begin == std::string::npos ? text : text.substr(begin); + const size_t end = res.find("<|"); + if (end != std::string::npos) { + res.resize(end); + } + for (auto & c : res) { + if (c == '\n' || c == '\r' || c == '\t') { + c = ' '; + } + } + const size_t first = res.find_first_not_of(' '); + res = first == std::string::npos ? "" : res.substr(first); + // ~80 characters + size_t bytes = 0; + for (size_t chars = 0; bytes < res.size() && chars < EXCERPT_CHARS; ++chars) { + bytes++; + while (bytes < res.size() && ((unsigned char) res[bytes] >> 6) == 0x2) { + bytes++; + } + } + return clean_utf8(res.substr(0, bytes), 4*EXCERPT_CHARS); +} + +int class_rank(cache_class c) { + return (int) c; +} + +bool expired(const cache_meta & m, int64_t now) { + return m.expires_at != 0 && m.expires_at <= now; +} + +// eviction order: expired, never-reused ephemeral then default (oldest first), then ephemeral < default < keep +// by LRU, pin last (only where it still has a tier below); false = may not be evicted. +// one-shot calls (hits == 0) go before any conversation that was continued, whatever their class. +bool evict_key(const cache_meta & m, int64_t now, bool allow_pin, std::tuple & key) { + if (expired(m, now)) { + key = { 0, 0, m.expires_at }; + } else if (m.cls == CACHE_CLASS_PIN) { + if (!allow_pin) { + return false; + } + key = { 4, 0, m.last_used }; + } else if (m.hits == 0 && m.cls == CACHE_CLASS_EPHEMERAL) { + key = { 1, 0, m.created }; + } else if (m.hits == 0 && m.cls == CACHE_CLASS_DEFAULT) { + key = { 2, 0, m.created }; + } else { + key = { 3, class_rank(m.cls), m.last_used }; + } + return true; +} + +// raises into the class of an entry it replaces (a shorter prompt of the same conversation) +bool merge_meta(cache_meta & into, const cache_meta & from) { + bool changed = false; + if (from.source == CACHE_SOURCE_EXPLICIT && + (into.source != CACHE_SOURCE_EXPLICIT || class_rank(from.cls) > class_rank(into.cls))) { + into.cls = std::max(into.source == CACHE_SOURCE_EXPLICIT ? into.cls : CACHE_CLASS_SKIP, from.cls); + into.source = CACHE_SOURCE_EXPLICIT; + if (into.ttl < 0) { + into.ttl = from.ttl; + } + changed = true; + } + if (from.hits > into.hits) { + into.hits = from.hits; + changed = true; + } + if (from.created != 0 && from.created < into.created) { + into.created = from.created; + changed = true; + } + if (into.excerpt.empty() && !from.excerpt.empty()) { + into.excerpt = from.excerpt; + changed = true; + } + into.family_saved = into.family_saved || from.family_saved; + into.family_reused = into.family_reused || from.family_reused; + return changed; +} + +template void put_le(std::vector & b, size_t off, T v) { + memcpy(b.data() + off, &v, sizeof(v)); +} + +template T get_le(const std::vector & b, size_t off) { + T v; + memcpy(&v, b.data() + off, sizeof(v)); + return v; +} + +std::vector meta_encode(const cache_meta & m) { + std::vector b(META_SIZE, 0); + put_le(b, 0, m.id); + put_le (b, 8, m.cls); + put_le (b, 9, m.source); + put_le (b, 10, m.family_saved); + put_le (b, 11, m.family_reused); + put_le (b, 16, m.ttl); + put_le (b, 24, m.created); + put_le (b, 32, m.last_used); + put_le (b, 40, m.expires_at); + put_le(b, 48, m.hits); + put_le(b, 56, m.family); + put_le (b, 64, m.n_input); + const std::string ex = clean_utf8(m.excerpt, META_SIZE - META_FIXED); + put_le(b, 72, (uint16_t) ex.size()); + memcpy(b.data() + META_FIXED, ex.data(), ex.size()); + return b; +} + +bool meta_decode(const std::vector & b, cache_meta & m) { + if (b.size() != META_SIZE || get_le(b, 8) > CACHE_CLASS_PIN || get_le(b, 9) > CACHE_SOURCE_EXPLICIT) { + return false; + } + m.id = get_le(b, 0); + m.cls = (cache_class) get_le(b, 8); + m.source = (cache_source) get_le(b, 9); + m.family_saved = get_le(b, 10) != 0; + m.family_reused = get_le(b, 11) != 0; + m.ttl = get_le(b, 16); + m.created = get_le(b, 24); + m.last_used = get_le(b, 32); + m.expires_at = get_le(b, 40); + m.hits = get_le(b, 48); + m.family = get_le(b, 56); + m.n_input = get_le(b, 64); + const size_t n = std::min(get_le(b, 72), META_SIZE - META_FIXED); + m.excerpt.assign((const char *) b.data() + META_FIXED, n); + return true; +} + struct disk_file { FILE * f; disk_file(const std::string & path, const char * mode) : f(fopen(path.c_str(), mode)) {} @@ -1743,22 +1992,40 @@ bool get_bytes(FILE * f, std::vector & v) { return get(f, v.data(), n); } -// magic, version, key and the prompt tokens: enough to index a file without reading its state -bool read_header(FILE * f, uint64_t key, llama_tokens & tokens) { +// magic, version, key, metadata (v2) and the prompt tokens: enough to index a file without reading its state. +// a v1 file has no metadata; has_meta is false for it +bool read_header(FILE * f, uint64_t key, llama_tokens & tokens, uint32_t & version, cache_meta & meta) { uint32_t magic; - uint32_t version; uint64_t k; uint64_t n; - if (!get(f, magic) || !get(f, version) || !get(f, k) || !get(f, n)) { + if (!get(f, magic) || !get(f, version) || !get(f, k)) { + return false; + } + if (magic != DISK_MAGIC || (version != 1 && version != DISK_VERSION) || k != key) { return false; } - if (magic != DISK_MAGIC || version != DISK_VERSION || k != key || n > (1ull << 26)) { + if (version == DISK_VERSION) { + std::vector b(META_SIZE); + if (!get(f, b.data(), META_SIZE) || !meta_decode(b, meta)) { + return false; + } + } + if (!get(f, n) || n > (1ull << 26)) { return false; } tokens.resize(n); return get(f, tokens.data(), n * sizeof(llama_token)); } +bool write_meta(const std::string & path, const cache_meta & meta) { + disk_file f(path, "r+b"); + if (!f.f || fseek(f.f, META_OFFSET, SEEK_SET) != 0) { + return false; + } + const auto b = meta_encode(meta); + return put(f.f, b.data(), b.size()) && fflush(f.f) == 0; +} + size_t common_prefix(const llama_tokens & a, const server_tokens & b) { const size_t n = std::min(a.size(), b.size()); size_t i = 0; @@ -1782,6 +2049,81 @@ bool is_prefix(const llama_tokens & a, const llama_tokens & b) { return a.size() <= b.size() && std::equal(a.begin(), a.end(), b.begin()); } +size_t count_seq(const llama_tokens & tokens, const llama_tokens & seq) { + if (seq.empty() || tokens.size() < seq.size()) { + return 0; + } + size_t n = 0; + for (size_t i = 0; i + seq.size() <= tokens.size(); ++i) { + if (std::equal(seq.begin(), seq.end(), tokens.begin() + i)) { + n++; + } + } + return n; +} + +json entry_json(const cache_meta & m, size_t n_tokens, size_t bytes, int64_t now) { + return { + { "id", hex_id(m.id) }, + { "class", cache_class_name(m.cls) }, + { "source", cache_source_name(m.source) }, + { "classified", m.classified() }, + { "n_tokens", n_tokens }, + { "bytes", bytes }, + { "hits", m.hits }, + { "age", now - m.created }, + { "idle", now - m.last_used }, + { "expires_in", m.expires_at == 0 ? json(nullptr) : json(m.expires_at - now) }, + { "family", hex_id(m.family) }, + { "excerpt", m.excerpt }, + }; +} + +// totals, per-class counts/bytes, pinned bytes and the newest INFO_ENTRIES entries +struct tier_report { + size_t prompts = 0; + size_t tokens = 0; + size_t bytes = 0; + size_t pinned = 0; + std::map> by_class; + std::vector> entries; + + void add(const cache_meta & m, size_t n_tokens, size_t n_bytes, int64_t now) { + prompts++; + tokens += n_tokens; + bytes += n_bytes; + if (m.cls == CACHE_CLASS_PIN) { + pinned += n_bytes; + } + auto & c = by_class[cache_class_name(m.cls)]; + c.first++; + c.second += n_bytes; + entries.push_back({ m.last_used, entry_json(m, n_tokens, n_bytes, now) }); + } + + void to_json(json & out) { + out["prompts"] = prompts; + out["tokens"] = tokens; + out["bytes"] = bytes; + out["pinned_bytes"] = pinned; + json classes = json::object(); + for (cache_class c : { CACHE_CLASS_EPHEMERAL, CACHE_CLASS_DEFAULT, CACHE_CLASS_KEEP, CACHE_CLASS_PIN }) { + const auto it = by_class.find(cache_class_name(c)); + classes[cache_class_name(c)] = { + { "count", it == by_class.end() ? 0 : it->second.first }, + { "bytes", it == by_class.end() ? 0 : it->second.second }, + }; + } + out["by_class"] = classes; + std::stable_sort(entries.begin(), entries.end(), [](const auto & a, const auto & b) { return a.first > b.first; }); + json list = json::array(); + for (size_t i = 0; i < entries.size() && i < INFO_ENTRIES; ++i) { + list.push_back(entries[i].second); + } + out["entries"] = list; + } +}; + } // namespace size_t server_prompt_cache::size() const { @@ -1805,12 +2147,17 @@ size_t server_prompt_cache::n_tokens() const { } server_prompt_cache_state * server_prompt_cache::alloc(const server_prompt & prompt, size_t state_size_tgt, size_t state_size_dft) { + expire_ram(); + // first check if the current state is contained fully in the cache for (auto it = states.begin(); it != states.end(); ++it) { const int cur_lcp_len = it->prompt.tokens.get_common_prefix(prompt.tokens); if (cur_lcp_len == (int) prompt.tokens.size()) { SRV_TRC("%s", " - prompt is already in the cache, skipping\n"); + if (merge_meta(it->prompt.meta, prompt.meta)) { + set_expiry(it->prompt.meta); + } return nullptr; } } @@ -1830,6 +2177,8 @@ server_prompt_cache_state * server_prompt_cache::alloc(const server_prompt & pro return nullptr; } + cache_meta meta = prompt.meta; + // remove any cached prompts that are fully contained in the current prompt for (auto it = states.begin(); it != states.end();) { const int len = it->prompt.tokens.get_common_prefix(prompt.tokens); @@ -1837,6 +2186,8 @@ server_prompt_cache_state * server_prompt_cache::alloc(const server_prompt & pro if (len == (int) it->prompt.tokens.size()) { SRV_TRC(" - removing obsolete cached prompt with length %d\n", len); + merge_meta(meta, it->prompt.meta); + forget_classify(it->prompt.meta.id); it = states.erase(it); } else { ++it; @@ -1846,10 +2197,15 @@ server_prompt_cache_state * server_prompt_cache::alloc(const server_prompt & pro if (limit_size > 0) { // make room before allocating the new vectors to avoid breaching the limit while (!states.empty() && size() + state_size_new > limit_size) { - SRV_WRN(" - making room for prompt cache entry, removing oldest entry (size = %.3f MiB)\n", - states.front().size() / (1024.0 * 1024.0)); + auto victim = ram_victim(disk_thread.joinable()); + if (victim == states.end()) { + SRV_WRN(" - prompt cache holds only pinned entries, not storing a prompt of %.3f MiB\n", state_size_new / (1024.0 * 1024.0)); + return nullptr; + } + SRV_WRN(" - making room for prompt cache entry, evicting a %s entry (size = %.3f MiB)\n", + cache_class_name(victim->prompt.meta.cls), victim->size() / (1024.0 * 1024.0)); - evict_front(); + evict(victim); } } @@ -1876,6 +2232,7 @@ server_prompt_cache_state * server_prompt_cache::alloc(const server_prompt & pro /*.prompt =*/ { /*.tokens =*/ prompt.tokens.clone(), /*.checkpoints =*/ prompt.checkpoints, + /*.meta =*/ std::move(meta), }, /*.data =*/ { /*.main =*/ std::move(state_data_tgt), @@ -1883,10 +2240,14 @@ server_prompt_cache_state * server_prompt_cache::alloc(const server_prompt & pro }, }); + note_saved(states.back()); + return &states.back(); } -bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_tgt, llama_context * ctx_dft, int32_t id_slot) { +bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_tgt, llama_context * ctx_dft, int32_t id_slot, bool keep_entry) { + expire_ram(); + const int lcp_best = prompt.tokens.get_common_prefix(tokens_new); float f_keep_best = prompt.tokens.size() > 0 ? float(lcp_best) / prompt.tokens.size() : -1.0f; // empty slot: any cache entry wins @@ -1923,6 +2284,8 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok if (disk_thread.joinable()) { std::lock_guard lock(disk_mutex); + expire_disk_locked(); + for (const auto & e : disk) { const int lcp_cur = common_prefix(e.tokens, tokens_new); @@ -1943,6 +2306,7 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok } } + bool from_disk = false; if (!disk_best.empty()) { namespace fs = std::filesystem; @@ -1956,12 +2320,21 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok if (ok) { SRV_INF("prompt cache: read %zu tokens from disk in %.2f s, f_keep = %.3f, f_sim = %.3f\n", state.prompt.tokens.size(), (ggml_time_us() - t_start) / 1e6, f_keep_best, f_sim_best); - states.push_back(std::move(state)); - it_best = std::prev(states.end()); - // most recently used: last to be evicted, here and after a restart if (it != disk.end()) { + // the file stays: it records the use, and is most recently used, here and after a restart + it->meta.last_used = unix_now(); + it->meta.hits++; + set_expiry(it->meta); + state.prompt.meta = it->meta; + state.prompt.meta.hits--; // the slot counts this use when the request continues the prompt + if (it->version == DISK_VERSION) { + write_meta(it->path, it->meta); + } disk.splice(disk.end(), disk, it); } + states.push_back(std::move(state)); + it_best = std::prev(states.end()); + from_disk = true; fs::last_write_time(disk_best, fs::file_time_type::clock::now(), ec); } else { SRV_WRN("prompt cache: cannot read %s, removing it\n", disk_best.c_str()); @@ -1975,6 +2348,9 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok if (it_best != states.end()) { SRV_TRC(" - found better prompt with f_keep = %.3f, f_sim = %.3f\n", f_keep_best, f_sim_best); + // a skip request restores a copy and leaves the RAM entry where it is + const bool consume = !keep_entry || from_disk; + { auto & data = it_best->data.main; @@ -1986,8 +2362,10 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok return false; } - data.clear(); - data.shrink_to_fit(); + if (consume) { + data.clear(); + data.shrink_to_fit(); + } } { @@ -2004,25 +2382,57 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok return false; } - data.clear(); - data.shrink_to_fit(); + if (consume) { + data.clear(); + data.shrink_to_fit(); + } } } - prompt = std::move(it_best->prompt); + if (consume) { + prompt = std::move(it_best->prompt); - states.erase(it_best); + states.erase(it_best); + } else { + prompt = it_best->prompt.clone(); + } } + update_ram_next_expiry(); + return true; } +std::list::iterator server_prompt_cache::ram_victim(bool allow_pin) { + const int64_t now = unix_now(); + auto best = states.end(); + std::tuple best_key; + for (auto it = states.begin(); it != states.end(); ++it) { + std::tuple key; + if (evict_key(it->prompt.meta, now, allow_pin, key) && (best == states.end() || key < best_key)) { + best = it; + best_key = key; + } + } + return best; +} + void server_prompt_cache::update() { + expire_ram(); + + const bool allow_pin = disk_thread.joinable(); + if (limit_size > 0) { while (!states.empty() && size() > limit_size) { - SRV_WRN(" - cache size limit reached, removing oldest entry (size = %.3f MiB)\n", states.front().size() / (1024.0 * 1024.0)); + auto victim = ram_victim(allow_pin); + if (victim == states.end()) { + SRV_WRN("%s", " - cache size limit reached, but only pinned entries are left\n"); + break; + } + SRV_WRN(" - cache size limit reached, evicting a %s entry (size = %.3f MiB)\n", + cache_class_name(victim->prompt.meta.cls), victim->size() / (1024.0 * 1024.0)); - evict_front(); + evict(victim); } } @@ -2034,13 +2444,19 @@ void server_prompt_cache::update() { if (limit_tokens > 0) { while (!states.empty() && n_tokens() > limit_tokens_cur) { - SRV_WRN(" - cache token limit (%zu, est: %zu) reached, removing oldest entry (size = %.3f MiB)\n", - limit_tokens, limit_tokens_cur, states.front().size() / (1024.0 * 1024.0)); + auto victim = ram_victim(allow_pin); + if (victim == states.end()) { + break; + } + SRV_WRN(" - cache token limit (%zu, est: %zu) reached, evicting a %s entry (size = %.3f MiB)\n", + limit_tokens, limit_tokens_cur, cache_class_name(victim->prompt.meta.cls), victim->size() / (1024.0 * 1024.0)); - evict_front(); + evict(victim); } } + update_ram_next_expiry(); + SRV_TRC(" - cache state: %zu prompts, %.3f MiB (limits: %.3f MiB, %zu tokens, %zu est)\n", states.size(), size() / (1024.0 * 1024.0), limit_size / (1024.0 * 1024.0), limit_tokens, limit_tokens_cur); @@ -2050,12 +2466,485 @@ void server_prompt_cache::update() { } } +void server_prompt_cache::expire_ram() { + const int64_t now = unix_now(); + for (auto it = states.begin(); it != states.end();) { + if (expired(it->prompt.meta, now)) { + SRV_INF("prompt cache: dropping expired %s entry %s (%d tokens)\n", + cache_class_name(it->prompt.meta.cls), hex_id(it->prompt.meta.id).c_str(), it->prompt.n_tokens()); + forget_classify(it->prompt.meta.id); + it = states.erase(it); + } else { + ++it; + } + } +} + +void server_prompt_cache::update_ram_next_expiry() { + int64_t next = 0; + for (const auto & s : states) { + const int64_t e = s.prompt.meta.expires_at; + if (e != 0 && (next == 0 || e < next)) { + next = e; + } + } + ram_next_expiry = next; +} + +// +// server_prompt_cache: retention policy +// + +void server_prompt_cache::policy_init(const cache_policy & policy, std::function detokenize, + llama_tokens user_marker, std::function post_maint) { + this->pol = policy; + this->detokenize = std::move(detokenize); + this->user_marker = std::move(user_marker); + this->post_maint = std::move(post_maint); +} + +void server_prompt_cache::policy_start() { + policy_thread = std::thread([this] { policy_loop(); }); +} + +void server_prompt_cache::set_expiry(cache_meta & m) const { + if (m.cls == CACHE_CLASS_PIN) { + m.expires_at = 0; + return; + } + const int64_t ttl = m.ttl >= 0 ? m.ttl : pol.ttl(m.cls); + m.expires_at = std::max(m.last_used, 1) + ttl; +} + +void server_prompt_cache::apply_request(cache_meta & m, bool continuation, const task_params & params, int64_t n_input) { + const int64_t now = unix_now(); + if (!continuation) { + m = {}; + m.created = now; + } else { + m.hits++; + auto f = families.find(m.family); + if (m.family_saved && !m.family_reused && f != families.end()) { + f->second.reused++; + f->second.last_seen = now; + families_dirty = true; + } + m.family_reused = m.family_reused || m.family_saved; + } + m.last_used = now; + m.n_input = n_input; + + if (params.cache_explicit && params.cache_cls != CACHE_CLASS_SKIP) { + // the most retentive explicit class wins + if (m.source != CACHE_SOURCE_EXPLICIT || class_rank(params.cache_cls) > class_rank(m.cls)) { + m.cls = params.cache_cls; + } + m.source = CACHE_SOURCE_EXPLICIT; + } + if (params.cache_ttl >= 0) { + m.ttl = params.cache_ttl; + } +} + +bool server_prompt_cache::family_class(uint64_t family, cache_class & out) const { + const auto it = families.find(family); + if (it == families.end() || it->second.saved < FAMILY_MIN_N) { + return false; + } + const double rate = double(it->second.reused) / it->second.saved; + out = rate < 0.05 ? CACHE_CLASS_EPHEMERAL : rate >= 0.30 ? CACHE_CLASS_KEEP : CACHE_CLASS_DEFAULT; + return true; +} + +void server_prompt_cache::apply_family(cache_meta & m) const { + cache_class cls; + if (m.source != CACHE_SOURCE_EXPLICIT && family_class(m.family, cls)) { + m.cls = cls; + m.source = CACHE_SOURCE_FAMILY; + } +} + +// a new RAM entry: its id, family, excerpt and expiry; queued for the classifier if nothing else decided +void server_prompt_cache::note_saved(server_prompt_cache_state & state) { + auto & m = state.prompt.meta; + const int64_t now = unix_now(); + const llama_tokens tokens = raw_tokens(state.prompt.tokens); + + m.id = new_entry_id(); + if (m.created == 0) { + m.created = now; + } + if (m.last_used == 0) { + m.last_used = now; + } + const size_t n_head = std::min(tokens.size(), CLASSIFY_HEAD); + const llama_tokens head(tokens.begin(), tokens.begin() + n_head); + std::string head_text; + if (detokenize && (m.excerpt.empty() || (m.source == CACHE_SOURCE_DEFAULT && !pol.classifier_url.empty()))) { + head_text = detokenize(head); + } + if (m.excerpt.empty()) { + m.excerpt = user_excerpt(head_text); + } + + if (m.family == 0) { + const size_t n = std::min(tokens.size(), FAMILY_TOKENS); + m.family = fnv1a(tokens.data(), n * sizeof(llama_token)); + } + auto & fam = families[m.family]; + if (!m.family_saved) { + fam.saved++; + m.family_saved = true; + } + fam.last_seen = now; + if (fam.excerpt.empty()) { + fam.excerpt = m.excerpt; + } + families_dirty = true; + + apply_family(m); + set_expiry(m); + + if (m.source == CACHE_SOURCE_DEFAULT && !pol.classifier_url.empty()) { + std::string excerpt = head_text; + if (tokens.size() > CLASSIFY_HEAD) { + const size_t n_tail = std::min(tokens.size() - CLASSIFY_HEAD, CLASSIFY_TAIL); + excerpt += "\n[...]\n" + detokenize(llama_tokens(tokens.end() - n_tail, tokens.end())); + } + std::lock_guard lock(policy_mutex); + if (classify_queue.size() >= CLASSIFY_QUEUE_MAX) { + classify_queue.pop_front(); + } + classify_queue.push_back({ m.id, clean_utf8(excerpt, excerpt.size()), (int64_t) tokens.size(), (int64_t) count_seq(tokens, user_marker) }); + policy_cv.notify_all(); + } +} + +void server_prompt_cache::forget_classify(uint64_t id) { + std::lock_guard lock(policy_mutex); + classify_queue.erase(std::remove_if(classify_queue.begin(), classify_queue.end(), + [id](const classify_job & j) { return j.id == id; }), classify_queue.end()); +} + +void server_prompt_cache::policy_loop() { + std::unique_lock lock(policy_mutex); + while (!policy_stop) { + policy_cv.wait_for(lock, std::chrono::seconds(30), [this] { return policy_stop || !classify_queue.empty(); }); + if (policy_stop) { + break; + } + + if (!classify_queue.empty()) { + // one request in flight at a time + const classify_job job = std::move(classify_queue.front()); + classify_queue.pop_front(); + lock.unlock(); + + std::string base = pol.classifier_url; + std::string path = "/"; + const size_t scheme = base.find("://"); + const size_t slash = base.find('/', scheme == std::string::npos ? 0 : scheme + 3); + if (slash != std::string::npos) { + path = base.substr(slash); + base.resize(slash); + } + classify_result res = { job.id, CACHE_CLASS_DEFAULT, -1 }; + bool ok = false; + try { + httplib::Client cli(base); + cli.set_connection_timeout(pol.classifier_timeout, 0); + cli.set_read_timeout(pol.classifier_timeout, 0); + cli.set_write_timeout(pol.classifier_timeout, 0); + const json body = { { "excerpt", job.excerpt }, { "n_tokens", job.n_tokens }, { "n_turns", job.n_turns } }; + auto r = cli.Post(path, body.dump_safe(), "application/json"); + if (r && r->status == 200) { + const json j = json::parse(r->body); + ok = j.contains("class") && j.at("class").is_string() && cache_class_from_name(j.at("class").get(), res.cls); + if (ok && j.contains("ttl") && j.at("ttl").is_number()) { + res.ttl = std::max(0, j.at("ttl").get()); + } + } + if (!ok) { + SRV_WRN("prompt cache: classifier gave no class for %s (%s)\n", hex_id(job.id).c_str(), + r ? string_format("status %d", r->status).c_str() : httplib::to_string(r.error()).c_str()); + } + } catch (const std::exception & e) { + SRV_WRN("prompt cache: classifier reply for %s unreadable: %s\n", hex_id(job.id).c_str(), e.what()); + } + + lock.lock(); + if (ok) { + classified.push_back(res); + } + } + + // disk expiry needs no server loop + if (disk_thread.joinable()) { + lock.unlock(); + { + std::lock_guard dlock(disk_mutex); + expire_disk_locked(); + } + lock.lock(); + } + + const int64_t next = ram_next_expiry; + const bool due = !classified.empty() || (next != 0 && next <= unix_now()); + if (due && !maint_posted && post_maint) { + maint_posted = true; + post_maint(); + } + } +} + +void server_prompt_cache::maintain(const std::vector & live) { + std::vector results; + { + std::lock_guard lock(policy_mutex); + results.swap(classified); + maint_posted = false; + } + + // classifier answers: never over an explicit class or a family verdict + for (const auto & r : results) { + auto apply = [&](cache_meta & m) { + if (m.source == CACHE_SOURCE_EXPLICIT || m.source == CACHE_SOURCE_FAMILY) { + return false; + } + m.cls = r.cls; + m.source = CACHE_SOURCE_CLASSIFIER; + if (r.ttl >= 0) { + m.ttl = r.ttl; + } + set_expiry(m); + return true; + }; + if (r.cls == CACHE_CLASS_SKIP) { + SRV_INF("prompt cache: classifier says skip, dropping %s\n", hex_id(r.id).c_str()); + delete_entry(r.id); + continue; + } + bool found = false; + for (auto & s : states) { + if (s.prompt.meta.id == r.id) { + apply(s.prompt.meta); + found = true; + } + } + for (auto * m : live) { + if (m->id == r.id) { + apply(*m); + found = true; + } + } + if (!found && disk_thread.joinable()) { + std::lock_guard lock(disk_mutex); + for (auto & s : pending) { + if (s.prompt.meta.id == r.id) { + apply(s.prompt.meta); + } + } + for (auto & e : disk) { + if (e.meta.id == r.id && apply(e.meta)) { + disk_rewrite_meta_locked(e); + } + } + } + SRV_INF("prompt cache: classifier: %s is %s\n", hex_id(r.id).c_str(), cache_class_name(r.cls)); + } + + // family verdicts follow the stats as they move + for (auto & s : states) { + apply_family(s.prompt.meta); + set_expiry(s.prompt.meta); + } + expire_ram(); + update_ram_next_expiry(); + + if (disk_thread.joinable()) { + std::lock_guard lock(disk_mutex); + for (auto & e : disk) { + const cache_meta before = e.meta; + apply_family(e.meta); + if (e.meta.cls != before.cls || e.meta.source != before.source) { + set_expiry(e.meta); + disk_rewrite_meta_locked(e); + } + } + expire_disk_locked(); + } + + families_save(); +} + +bool server_prompt_cache::set_entry(uint64_t id, cache_class cls, int64_t ttl) { + bool found = false; + auto apply = [&](cache_meta & m) { + m.cls = cls; + m.source = CACHE_SOURCE_EXPLICIT; + m.ttl = ttl; + m.last_used = std::max(m.last_used, unix_now()); + set_expiry(m); + found = true; + }; + for (auto & s : states) { + if (s.prompt.meta.id == id) { + apply(s.prompt.meta); + } + } + if (disk_thread.joinable()) { + std::lock_guard lock(disk_mutex); + for (auto & s : pending) { + if (s.prompt.meta.id == id) { + apply(s.prompt.meta); + } + } + for (auto & e : disk) { + if (e.meta.id == id) { + apply(e.meta); + disk_rewrite_meta_locked(e); + } + } + } + update_ram_next_expiry(); + return found; +} + +bool server_prompt_cache::delete_entry(uint64_t id) { + bool found = false; + for (auto it = states.begin(); it != states.end();) { + if (it->prompt.meta.id == id) { + it = states.erase(it); + found = true; + } else { + ++it; + } + } + if (disk_thread.joinable()) { + std::lock_guard lock(disk_mutex); + std::error_code ec; + for (auto it = disk.begin(); it != disk.end();) { + if (it->meta.id == id) { + std::filesystem::remove(it->path, ec); + it = disk.erase(it); + found = true; + } else { + ++it; + } + } + } + forget_classify(id); + update_ram_next_expiry(); + return found; +} + +size_t server_prompt_cache::clear_ram(bool include_pinned) { + size_t n = 0; + for (auto it = states.begin(); it != states.end();) { + if (include_pinned || it->prompt.meta.cls != CACHE_CLASS_PIN) { + forget_classify(it->prompt.meta.id); + it = states.erase(it); + n++; + } else { + ++it; + } + } + update_ram_next_expiry(); + return n; +} + +json server_prompt_cache::family_info() const { + std::vector> list; + for (const auto & [h, f] : families) { + list.push_back({ h, &f }); + } + std::sort(list.begin(), list.end(), [](const auto & a, const auto & b) { return a.second->last_seen > b.second->last_seen; }); + const int64_t now = unix_now(); + json res = json::array(); + for (size_t i = 0; i < list.size() && i < INFO_ENTRIES; ++i) { + const auto & f = *list[i].second; + cache_class cls; + const bool decided = family_class(list[i].first, cls); + res.push_back({ + { "family", hex_id(list[i].first) }, + { "n", f.saved }, + { "reused", f.reused }, + { "reuse_rate", f.saved > 0 ? double(f.reused) / f.saved : 0.0 }, + { "class", decided ? json(cache_class_name(cls)) : json(nullptr) }, + { "last_seen", now - f.last_seen }, + { "excerpt", f.excerpt }, + }); + } + return res; +} + +void server_prompt_cache::families_load() { + const std::string path = disk_dir + "/families.json"; + std::ifstream in(path); + if (!in) { + return; + } + try { + const json j = json::parse(std::string(std::istreambuf_iterator(in), std::istreambuf_iterator())); + for (const auto & f : j) { + family_stat s; + s.saved = f.at("saved").get(); + s.reused = f.at("reused").get(); + s.last_seen = f.at("last_seen").get(); + s.excerpt = f.contains("excerpt") ? f.at("excerpt").get() : ""; + families[std::stoull(f.at("family").get(), nullptr, 16)] = s; + } + } catch (const std::exception & e) { + SRV_WRN("prompt cache: ignoring %s: %s\n", path.c_str(), e.what()); + families.clear(); + } +} + +// keeps the FAMILIES_MAX most recently seen families +void server_prompt_cache::families_save() { + if (families.size() > FAMILIES_MAX) { + std::vector seen; + for (const auto & [_, f] : families) { + seen.push_back(f.last_seen); + } + std::nth_element(seen.begin(), seen.begin() + (seen.size() - FAMILIES_MAX), seen.end()); + const int64_t cut = seen[seen.size() - FAMILIES_MAX]; + for (auto it = families.begin(); it != families.end() && families.size() > FAMILIES_MAX;) { + it = it->second.last_seen < cut ? families.erase(it) : std::next(it); + } + families_dirty = true; + } + if (!families_dirty || disk_dir.empty() || !disk_thread.joinable()) { + return; + } + json j = json::array(); + for (const auto & [h, f] : families) { + j.push_back({ { "family", hex_id(h) }, { "saved", f.saved }, { "reused", f.reused }, { "last_seen", f.last_seen }, { "excerpt", f.excerpt } }); + } + const std::string path = disk_dir + "/families.json"; + { + std::ofstream out(path + ".tmp"); + out << j.dump_safe(); + } + std::error_code ec; + std::filesystem::rename(path + ".tmp", path, ec); + families_dirty = false; +} + // // server_prompt_cache: disk tier // server_prompt_cache::~server_prompt_cache() { + if (policy_thread.joinable()) { + { + std::lock_guard lock(policy_mutex); + policy_stop = true; + } + policy_cv.notify_all(); + policy_thread.join(); + } if (disk_thread.joinable()) { { std::lock_guard lock(disk_mutex); @@ -2080,6 +2969,9 @@ void server_prompt_cache::disk_init(const std::string & dir, int32_t limit_mib, return; } + families_load(); + + size_t n_v1 = 0; std::vector> found; for (const auto & de : fs::directory_iterator(dir, ec)) { const auto ext = de.path().extension(); @@ -2091,103 +2983,118 @@ void server_prompt_cache::disk_init(const std::string & dir, int32_t limit_mib, continue; } llama_tokens tokens; + uint32_t version = 0; + cache_meta meta; bool ok; { disk_file f(de.path().string(), "rb"); - ok = f.f && read_header(f.f, disk_key, tokens); + ok = f.f && read_header(f.f, disk_key, tokens, version, meta); } if (!ok) { // another model, cache layout or format fs::remove(de.path(), ec); continue; } - found.push_back({ de.last_write_time(ec), { de.path().string(), std::move(tokens), (size_t) de.file_size(ec) } }); + const auto mtime = de.last_write_time(ec); + if (version == 1) { + // no metadata: class default, aged from the file's mtime + const auto sys = std::chrono::system_clock::now() + std::chrono::duration_cast( + mtime - fs::file_time_type::clock::now()); + const int64_t t = std::chrono::duration_cast(sys.time_since_epoch()).count(); + meta.id = new_entry_id(); + meta.created = t; + meta.last_used = t; + meta.family = fnv1a(tokens.data(), std::min(tokens.size(), FAMILY_TOKENS) * sizeof(llama_token)); + if (detokenize) { + meta.excerpt = user_excerpt(detokenize(llama_tokens(tokens.begin(), tokens.begin() + std::min(tokens.size(), CLASSIFY_HEAD)))); + } + set_expiry(meta); + n_v1++; + } + found.push_back({ mtime, { de.path().string(), std::move(tokens), (size_t) de.file_size(ec), std::move(meta), version } }); } std::sort(found.begin(), found.end(), [](const auto & a, const auto & b) { return a.first < b.first; }); - size_t total = 0; for (auto & [_, e] : found) { - total += e.size; disk.push_back(std::move(e)); } - while (disk_limit > 0 && total > disk_limit && !disk.empty()) { - total -= disk.front().size; - fs::remove(disk.front().path, ec); - disk.pop_front(); + expire_disk_locked(); + disk_evict_locked(""); + + size_t total = 0; + for (const auto & e : disk) { + total += e.size; } - SRV_INF("prompt cache: disk tier %s, %zu prompts, %.1f GiB (limit %.1f GiB)\n", - dir.c_str(), disk.size(), total / (1024.0*1024.0*1024.0), disk_limit / (1024.0*1024.0*1024.0)); + SRV_INF("prompt cache: disk tier %s, %zu prompts (%zu v1), %.1f GiB (limit %.1f GiB)\n", + dir.c_str(), disk.size(), n_v1, total / (1024.0*1024.0*1024.0), disk_limit / (1024.0*1024.0*1024.0)); disk_thread = std::thread([this] { disk_write_loop(); }); } -void server_prompt_cache::evict_front() { - if (disk_thread.joinable()) { +void server_prompt_cache::evict(std::list::iterator it) { + if (disk_thread.joinable() && !expired(it->prompt.meta, unix_now())) { // text prompts only: media chunks are not serialized - const llama_tokens tokens = raw_tokens(states.front().prompt.tokens); + const llama_tokens tokens = raw_tokens(it->prompt.tokens); if (std::find(tokens.begin(), tokens.end(), LLAMA_TOKEN_NULL) == tokens.end()) { std::lock_guard lock(disk_mutex); - pending.splice(pending.end(), states, states.begin()); + pending.splice(pending.end(), states, it); disk_cv.notify_all(); return; } } - states.pop_front(); + forget_classify(it->prompt.meta.id); + states.erase(it); } void server_prompt_cache::flush() { + families_save(); if (!disk_thread.joinable()) { return; } while (!states.empty()) { - evict_front(); + evict(states.begin()); } std::unique_lock lock(disk_mutex); - disk_cv.wait(lock, [this] { return pending.empty(); }); + disk_cv.wait(lock, [this] { return pending.empty() && upgrades.empty(); }); } void server_prompt_cache::disk_write_loop() { std::unique_lock lock(disk_mutex); while (true) { - disk_cv.wait(lock, [this] { return disk_stop || !pending.empty(); }); - if (pending.empty()) { + disk_cv.wait(lock, [this] { return disk_stop || !pending.empty() || !upgrades.empty(); }); + if (!pending.empty()) { + auto & state = pending.front(); + lock.unlock(); + disk_write(state); + lock.lock(); + pending.pop_front(); + } else if (!upgrades.empty()) { + const std::string path = upgrades.front(); + lock.unlock(); + disk_upgrade(path); + lock.lock(); + upgrades.pop_front(); + } else { return; } - auto & state = pending.front(); - lock.unlock(); - disk_write(state); - lock.lock(); - pending.pop_front(); disk_cv.notify_all(); } } -void server_prompt_cache::disk_write(server_prompt_cache_state & state) { - namespace fs = std::filesystem; - - const llama_tokens tokens = raw_tokens(state.prompt.tokens); - { - std::lock_guard lock(disk_mutex); - for (const auto & e : disk) { - if (is_prefix(tokens, e.tokens)) { - return; // a file already holds this prompt or a longer one - } - } - } - - const int64_t t_start = ggml_time_us(); - const uint64_t h = fnv1a(tokens.data(), tokens.size() * sizeof(llama_token)); - const std::string path = string_format("%s/%016llx-%zu.kvc", disk_dir.c_str(), (unsigned long long) h, tokens.size()); - const std::string tmp = path + ".tmp"; +namespace { +// the whole file at path (via path.tmp): header, metadata, then the v1 body +bool write_file(const std::string & path, uint64_t key, const llama_tokens & tokens, const server_prompt_cache_state & state, const cache_meta & meta) { + const std::string tmp = path + ".tmp"; bool ok; { disk_file f(tmp, "wb"); ok = f.f != nullptr; if (ok) { setvbuf(f.f, nullptr, _IOFBF, 8u << 20); - ok = put(f.f, DISK_MAGIC) && put(f.f, DISK_VERSION) && put(f.f, disk_key) && + const auto mb = meta_encode(meta); + ok = put(f.f, DISK_MAGIC) && put(f.f, DISK_VERSION) && put(f.f, key) && put(f.f, mb.data(), mb.size()) && put(f.f, (uint64_t) tokens.size()) && put(f.f, tokens.data(), tokens.size() * sizeof(llama_token)) && put(f.f, (uint64_t) state.prompt.checkpoints.size()); for (const auto & ckpt : state.prompt.checkpoints) { @@ -2200,21 +3107,156 @@ void server_prompt_cache::disk_write(server_prompt_cache_state & state) { } std::error_code ec; if (ok) { - fs::rename(tmp, path, ec); + std::filesystem::rename(tmp, path, ec); ok = !ec; } if (!ok) { + std::filesystem::remove(tmp, ec); + } + return ok; +} + +} // namespace + +void server_prompt_cache::disk_rewrite_meta_locked(disk_entry & e) { + if (e.version == DISK_VERSION) { + if (!write_meta(e.path, e.meta)) { + SRV_WRN("prompt cache: cannot update metadata of %s\n", e.path.c_str()); + } + } else if (std::find(upgrades.begin(), upgrades.end(), e.path) == upgrades.end()) { + upgrades.push_back(e.path); + disk_cv.notify_all(); + } +} + +// rewrites a v1 file as v2 so its metadata can be stored +void server_prompt_cache::disk_upgrade(const std::string & path) { + server_prompt_cache_state state; + if (!disk_read(path, false, state)) { + return; + } + cache_meta meta; + { + std::lock_guard lock(disk_mutex); + auto it = std::find_if(disk.begin(), disk.end(), [&](const disk_entry & e) { return e.path == path; }); + if (it == disk.end() || it->version == DISK_VERSION) { + return; + } + meta = it->meta; + } + const llama_tokens tokens = raw_tokens(state.prompt.tokens); + const bool ok = write_file(path, disk_key, tokens, state, meta); + + std::lock_guard lock(disk_mutex); + auto it = std::find_if(disk.begin(), disk.end(), [&](const disk_entry & e) { return e.path == path; }); + if (it == disk.end()) { + return; + } + if (!ok) { + SRV_ERR("prompt cache: failed to upgrade %s\n", path.c_str()); + return; + } + std::error_code ec; + it->version = DISK_VERSION; + it->size = std::filesystem::file_size(path, ec); + if (meta_encode(it->meta) != meta_encode(meta)) { + write_meta(path, it->meta); + } + SRV_INF("prompt cache: upgraded %s to v%u\n", path.c_str(), DISK_VERSION); +} + +void server_prompt_cache::expire_disk_locked() { + const int64_t now = unix_now(); + std::error_code ec; + for (auto it = disk.begin(); it != disk.end();) { + if (expired(it->meta, now)) { + SRV_INF("prompt cache: removing expired %s file %s\n", cache_class_name(it->meta.cls), it->path.c_str()); + std::filesystem::remove(it->path, ec); + it = disk.erase(it); + } else { + ++it; + } + } +} + +// enforces the byte limit by the eviction order; pins stay, and a pinned newcomer that does not fit is dropped +void server_prompt_cache::disk_evict_locked(const std::string & newest) { + if (disk_limit == 0) { + return; + } + size_t total = 0; + for (const auto & e : disk) { + total += e.size; + } + const int64_t now = unix_now(); + std::error_code ec; + while (total > disk_limit) { + auto victim = disk.end(); + std::tuple best_key; + for (auto it = disk.begin(); it != disk.end(); ++it) { + std::tuple key; + if (evict_key(it->meta, now, false, key) && (victim == disk.end() || key < best_key)) { + victim = it; + best_key = key; + } + } + if (victim == disk.end()) { + victim = std::find_if(disk.begin(), disk.end(), [&](const disk_entry & e) { return e.path == newest; }); + if (victim == disk.end()) { + SRV_WRN("prompt cache: disk tier over its limit with pinned entries only (%.1f GiB)\n", total / (1024.0*1024.0*1024.0)); + return; + } + SRV_WRN("prompt cache: disk tier holds only pinned entries, not storing %s\n", newest.c_str()); + } else { + SRV_INF("prompt cache: evicting %s file %s (hits %llu)\n", cache_class_name(victim->meta.cls), victim->path.c_str(), + (unsigned long long) victim->meta.hits); + } + total -= victim->size; + std::filesystem::remove(victim->path, ec); + disk.erase(victim); + } +} + +void server_prompt_cache::disk_write(server_prompt_cache_state & state) { + namespace fs = std::filesystem; + + const llama_tokens tokens = raw_tokens(state.prompt.tokens); + cache_meta meta; + { + std::lock_guard lock(disk_mutex); + for (auto & e : disk) { + if (is_prefix(tokens, e.tokens)) { + // a file already holds this prompt or a longer one + if (merge_meta(e.meta, state.prompt.meta)) { + set_expiry(e.meta); + disk_rewrite_meta_locked(e); + } + return; + } + } + meta = state.prompt.meta; + } + + const int64_t t_start = ggml_time_us(); + const uint64_t h = fnv1a(tokens.data(), tokens.size() * sizeof(llama_token)); + const std::string path = string_format("%s/%016llx-%zu.kvc", disk_dir.c_str(), (unsigned long long) h, tokens.size()); + + if (!write_file(path, disk_key, tokens, state, meta)) { SRV_ERR("prompt cache: failed to write %s\n", path.c_str()); - fs::remove(tmp, ec); return; } + std::error_code ec; const size_t size = fs::file_size(path, ec); std::lock_guard lock(disk_mutex); + // changed (classifier, entry api) while it was being written + cache_meta latest = state.prompt.meta; + // files whose prompt this one extends are obsolete for (auto it = disk.begin(); it != disk.end();) { if (is_prefix(it->tokens, tokens) && it->path != path) { + merge_meta(latest, it->meta); fs::remove(it->path, ec); it = disk.erase(it); } else { @@ -2222,63 +3264,87 @@ void server_prompt_cache::disk_write(server_prompt_cache_state & state) { } } disk.remove_if([&](const disk_entry & e) { return e.path == path; }); - disk.push_back({ path, tokens, size }); + disk.push_back({ path, tokens, size, latest, DISK_VERSION }); + if (meta_encode(latest) != meta_encode(meta)) { + set_expiry(disk.back().meta); + write_meta(path, disk.back().meta); + } + + disk_evict_locked(path); size_t total = 0; for (const auto & e : disk) { total += e.size; } - while (disk_limit > 0 && total > disk_limit && disk.size() > 1) { - total -= disk.front().size; - fs::remove(disk.front().path, ec); - disk.pop_front(); - } - SRV_INF("prompt cache: wrote %zu tokens to disk, %.1f MiB in %.2f s (%zu prompts, %.1f GiB)\n", - tokens.size(), size / (1024.0*1024.0), (ggml_time_us() - t_start) / 1e6, disk.size(), total / (1024.0*1024.0*1024.0)); + SRV_INF("prompt cache: wrote %zu tokens (%s, %s) to disk, %.1f MiB in %.2f s (%zu prompts, %.1f GiB)\n", + tokens.size(), hex_id(meta.id).c_str(), cache_class_name(meta.cls), size / (1024.0*1024.0), (ggml_time_us() - t_start) / 1e6, + disk.size(), total / (1024.0*1024.0*1024.0)); } -json server_prompt_cache::info() const { +json server_prompt_cache::info() { + const int64_t now = unix_now(); + + json ram = { { "limit_bytes", limit_size } }; + { + tier_report rep; + for (const auto & s : states) { + rep.add(s.prompt.meta, s.prompt.tokens.size(), s.size(), now); + } + rep.to_json(ram); + } json res = { - { "ram", { - { "prompts", states.size() }, - { "tokens", n_tokens() }, - { "bytes", size() }, - { "limit_bytes", limit_size }, + { "ram", ram }, + { "families", family_info() }, + { "policy", { + { "ttl_ephemeral", pol.ttl_ephemeral }, + { "ttl_default", pol.ttl_default }, + { "ttl_keep", pol.ttl_keep }, + { "classifier", pol.classifier_url }, + { "classifier_timeout", pol.classifier_timeout }, } }, }; + { + std::lock_guard lock(policy_mutex); + res["policy"]["classifier_queue"] = classify_queue.size(); + } if (disk_thread.joinable()) { std::lock_guard lock(disk_mutex); - size_t bytes = 0; - size_t tokens = 0; + tier_report rep; + size_t n_v1 = 0; for (const auto & e : disk) { - bytes += e.size; - tokens += e.tokens.size(); + rep.add(e.meta, e.tokens.size(), e.size, now); + n_v1 += e.version == 1; } - res["disk"] = { + json d = { { "dir", disk_dir }, - { "prompts", disk.size() }, - { "tokens", tokens }, - { "bytes", bytes }, { "limit_bytes", disk_limit }, { "pending", pending.size() }, + { "v1_files", n_v1 }, }; + rep.to_json(d); + res["disk"] = d; } return res; } -size_t server_prompt_cache::clear_disk() { +size_t server_prompt_cache::clear_disk(bool include_pinned) { if (!disk_thread.joinable()) { return 0; } std::unique_lock lock(disk_mutex); - disk_cv.wait(lock, [this] { return pending.empty(); }); + disk_cv.wait(lock, [this] { return pending.empty() && upgrades.empty(); }); std::error_code ec; - for (const auto & e : disk) { - std::filesystem::remove(e.path, ec); + size_t n = 0; + for (auto it = disk.begin(); it != disk.end();) { + if (include_pinned || it->meta.cls != CACHE_CLASS_PIN) { + std::filesystem::remove(it->path, ec); + it = disk.erase(it); + n++; + } else { + ++it; + } } - const size_t n = disk.size(); - disk.clear(); SRV_INF("prompt cache: cleared %zu prompts from disk\n", n); return n; } @@ -2291,8 +3357,10 @@ bool server_prompt_cache::disk_read(const std::string & path, bool has_mtmd, ser setvbuf(f.f, nullptr, _IOFBF, 8u << 20); llama_tokens tokens; + uint32_t version; + cache_meta meta; uint64_t n_ckpt; - if (!read_header(f.f, disk_key, tokens) || !get(f.f, n_ckpt) || n_ckpt > 4096) { + if (!read_header(f.f, disk_key, tokens, version, meta) || !get(f.f, n_ckpt) || n_ckpt > 4096) { return false; } std::list checkpoints; @@ -2312,6 +3380,7 @@ bool server_prompt_cache::disk_read(const std::string & path, bool has_mtmd, ser out.prompt.tokens = server_tokens(tokens, has_mtmd); out.prompt.checkpoints = std::move(checkpoints); + out.prompt.meta = std::move(meta); out.data = std::move(data); return true; } diff --git a/tools/server/server-task.h b/tools/server/server-task.h index eb69e521..fcf7ae34 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -10,6 +10,9 @@ #include #include #include +#include +#include +#include // TODO: prevent including the whole server-common.h as we only use server_tokens #include "server-common.h" @@ -30,6 +33,8 @@ enum server_task_type { SERVER_TASK_TYPE_SLOT_ERASE, SERVER_TASK_TYPE_CACHE_GET, SERVER_TASK_TYPE_CACHE_CLEAR, + SERVER_TASK_TYPE_CACHE_ENTRY, + SERVER_TASK_TYPE_CACHE_MAINT, SERVER_TASK_TYPE_GET_LORA, SERVER_TASK_TYPE_SET_LORA, }; @@ -45,6 +50,27 @@ enum task_response_type { TASK_RESPONSE_TYPE_ANTHROPIC, }; +// retention class of a prompt cache entry, least retentive first +enum cache_class : uint8_t { + CACHE_CLASS_SKIP, + CACHE_CLASS_EPHEMERAL, + CACHE_CLASS_DEFAULT, + CACHE_CLASS_KEEP, + CACHE_CLASS_PIN, +}; + +// where an entry's class came from, lowest precedence first +enum cache_source : uint8_t { + CACHE_SOURCE_DEFAULT, + CACHE_SOURCE_CLASSIFIER, + CACHE_SOURCE_FAMILY, + CACHE_SOURCE_EXPLICIT, +}; + +const char * cache_class_name(cache_class cls); +bool cache_class_from_name(const std::string & name, cache_class & out); +const char * cache_source_name(cache_source src); + enum stop_type { STOP_TYPE_NONE, STOP_TYPE_EOS, @@ -56,6 +82,11 @@ struct task_params { bool stream = false; bool include_usage = false; bool cache_prompt = true; // remember the prompt to avoid reprocessing all prompt + + // "cache" and "cache_ttl": retention of the prompt this request leaves in its slot + cache_class cache_cls = CACHE_CLASS_DEFAULT; + bool cache_explicit = false; // "cache" was given + int64_t cache_ttl = -1; // seconds, -1 = the class default bool return_tokens = false; bool return_progress = false; @@ -179,6 +210,13 @@ struct server_task { // used by SERVER_TASK_TYPE_CACHE_CLEAR: "slots", "ram", "disk" or "all" std::string cache_scope; + bool cache_include_pinned = false; + + // used by SERVER_TASK_TYPE_CACHE_ENTRY: set the class (and ttl) of one entry, or delete it + std::string cache_entry_id; + std::string cache_entry_class; + std::string cache_entry_action; + int64_t cache_entry_ttl = -1; // used by SERVER_TASK_TYPE_SET_LORA std::map set_lora; // mapping adapter ID -> scale @@ -580,14 +618,38 @@ struct server_task_result_apply_lora : server_task_result { virtual json to_json() override; }; +// retention metadata of a prompt; it travels with the prompt between the cache and a slot +struct cache_meta { + uint64_t id = 0; // fresh for every stored entry + cache_class cls = CACHE_CLASS_DEFAULT; + cache_source source = CACHE_SOURCE_DEFAULT; + int64_t ttl = -1; // seconds, -1 = the class default + int64_t created = 0; // unix seconds + int64_t last_used = 0; + int64_t expires_at = 0; // 0 = never + uint64_t hits = 0; // times a later request continued this prompt + uint64_t family = 0; // hash of the first 256 tokens + int64_t n_input = 0; // prompt tokens of the request that produced the state + bool family_saved = false; // counted in its family's saved/reused stats + bool family_reused = false; + std::string excerpt; // start of the first user message + + bool classified() const { + return source != CACHE_SOURCE_DEFAULT; + } +}; + struct server_prompt { server_tokens tokens; std::list checkpoints; + cache_meta meta; + void clear() { tokens.clear(); checkpoints.clear(); + meta = {}; } int n_tokens() const { @@ -598,6 +660,7 @@ struct server_prompt { return server_prompt { tokens.clone(), checkpoints, + meta, }; } }; @@ -626,6 +689,17 @@ struct server_prompt_cache_state { } }; +struct cache_policy { + int64_t ttl_ephemeral = 3600; + int64_t ttl_default = 3*86400; + int64_t ttl_keep = 30*86400; + + std::string classifier_url; // empty = off + int32_t classifier_timeout = 20; // seconds + + int64_t ttl(cache_class cls) const; +}; + struct server_prompt_cache { server_prompt_cache(int32_t limit_size_mib, size_t limit_tokens) { this->limit_size = 1024ull*1024ull*(limit_size_mib < 0 ? 0 : limit_size_mib); @@ -648,7 +722,8 @@ struct server_prompt_cache { server_prompt_cache_state * alloc(const server_prompt & prompt, size_t state_size_main, size_t state_size_drft); - bool load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_tgt, llama_context * ctx_dft, int32_t id_slot); + // keep_entry: restore from a RAM entry without consuming it (a "skip" request) + bool load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_tgt, llama_context * ctx_dft, int32_t id_slot, bool keep_entry = false); void update(); @@ -659,36 +734,122 @@ struct server_prompt_cache { void flush(); - // sizes and limits of both tiers - json info() const; + // sizes, per-class totals, entries and families of both tiers + json info(); - // drops every disk entry and its file, after the writer has finished what is queued; returns the count - size_t clear_disk(); + // drops every disk entry and its file (pinned ones only with include_pinned), after the writer has + // finished what is queued; returns the count + size_t clear_disk(bool include_pinned = true); + + // drops RAM entries (pinned ones only with include_pinned); returns the count + size_t clear_ram(bool include_pinned = true); + + // + // retention policy + // + + // set before disk_init; detokenize renders special tokens, user_marker starts a user turn + void policy_init(const cache_policy & policy, std::function detokenize, + llama_tokens user_marker, std::function post_maint); + + // starts the classifier and expiry thread; after disk_init + void policy_start(); + + // applies a request's cache flags to the meta of the slot prompt it continues (or starts) + void apply_request(cache_meta & meta, bool continuation, const task_params & params, int64_t n_input); + + // classifier answers, expiry and family verdicts; live holds the metas of the slots' prompts + void maintain(const std::vector & live); + + // false if no entry has this id + bool set_entry(uint64_t id, cache_class cls, int64_t ttl); + bool delete_entry(uint64_t id); + + const cache_policy & policy() const { return pol; } private: struct disk_entry { std::string path; llama_tokens tokens; size_t size; + cache_meta meta; + uint32_t version; + }; + + struct family_stat { + uint64_t saved = 0; + uint64_t reused = 0; + int64_t last_seen = 0; + std::string excerpt; + }; + + struct classify_job { + uint64_t id; + std::string excerpt; + int64_t n_tokens; + int64_t n_turns; + }; + + struct classify_result { + uint64_t id; + cache_class cls; + int64_t ttl; }; std::string disk_dir; size_t disk_limit = 0; uint64_t disk_key = 0; - // oldest first; both lists and stop are guarded by disk_mutex + // oldest first; the lists and stop are guarded by disk_mutex std::list disk; std::list pending; + std::list upgrades; // v1 files to rewrite as v2 for their metadata mutable std::mutex disk_mutex; std::condition_variable disk_cv; std::thread disk_thread; bool disk_stop = false; - void evict_front(); + cache_policy pol; + std::function detokenize; + llama_tokens user_marker; + std::function post_maint; + + // server loop only + std::map families; + bool families_dirty = false; + + // guarded by policy_mutex + std::deque classify_queue; + std::vector classified; + std::atomic ram_next_expiry { 0 }; + bool maint_posted = false; + bool policy_stop = false; + std::mutex policy_mutex; + std::condition_variable policy_cv; + std::thread policy_thread; + + void evict(std::list::iterator it); + std::list::iterator ram_victim(bool allow_pin); + void expire_ram(); + void expire_disk_locked(); void disk_write_loop(); void disk_write(server_prompt_cache_state & state); + void disk_upgrade(const std::string & path); bool disk_read(const std::string & path, bool has_mtmd, server_prompt_cache_state & out) const; + void disk_evict_locked(const std::string & newest); + void disk_rewrite_meta_locked(disk_entry & e); + + void policy_loop(); + bool family_class(uint64_t family, cache_class & out) const; + void apply_family(cache_meta & meta) const; + void set_expiry(cache_meta & meta) const; + void note_saved(server_prompt_cache_state & state); + void forget_classify(uint64_t id); + void families_load(); + void families_save(); + void update_ram_next_expiry(); + json family_info() const; }; // used exclusively by router mode diff --git a/tools/server/server.cpp b/tools/server/server.cpp index 3bb19082..dc13a51c 100644 --- a/tools/server/server.cpp +++ b/tools/server/server.cpp @@ -221,6 +221,7 @@ int llama_server(common_params & params, int argc, char ** argv) { routes.post_slots = models_routes->proxy_post; routes.get_cache = models_routes->proxy_get; routes.post_cache_clear = models_routes->proxy_post; + routes.post_cache_entry = models_routes->proxy_post; // custom routes for router routes.get_props = models_routes->get_router_props; @@ -276,6 +277,7 @@ int llama_server(common_params & params, int argc, char ** argv) { ctx_http.post("/slots/:id_slot", ex_wrapper(routes.post_slots)); ctx_http.get ("/cache", ex_wrapper(routes.get_cache)); ctx_http.post("/cache/clear", ex_wrapper(routes.post_cache_clear)); + ctx_http.post("/cache/entry", ex_wrapper(routes.post_cache_entry)); // resumable streaming: a child binds the local session factories, the router binds // proxies that resolve the owning child, see server-stream.h diff --git a/tools/server/tests/cache_policy_test.py b/tools/server/tests/cache_policy_test.py new file mode 100644 index 00000000..1b24d804 --- /dev/null +++ b/tools/server/tests/cache_policy_test.py @@ -0,0 +1,329 @@ +#!/usr/bin/env python3 +"""cache retention policy against the real model on valefar; run through gpu.sh: + + /home/regent/bonsai-bench/gpu.sh python3 tools/server/tests/cache_policy_test.py BIN + +three server runs on 127.0.0.1:8898 with a /tmp cache dir: (1) skip, greedy, ttl, classifier, family, +entry api, clear, (2) restart: pin and metadata, v1 files, dead classifier, (3) eviction order. +""" +import glob, json, os, shutil, struct, subprocess, sys, threading, time, urllib.request, urllib.error +from http.server import BaseHTTPRequestHandler, HTTPServer + +BIN = sys.argv[1] if len(sys.argv) > 1 else "/home/regent/bonsai-cachepol/build/bin/llama-server" +M = "/home/regent/models/bonsai2-27b/Ternary-Bonsai-2-27B-PTQ1_0-mtp-valefar.gguf" +DIR = "/tmp/cachepol-test" +LOG = "/tmp/cachepol-test-logs" +URL = "http://127.0.0.1:8898" +rows = sorted((json.loads(l) for l in open("/home/regent/bonsai-mi/data/heldout-long.jsonl")), key=lambda x: -len(json.dumps(x))) +FILLER = " ".join(m["content"] for r in rows for m in r["messages"] if isinstance(m.get("content"), str)) +results = [] +proc = None + + +def check(name, ok, detail=""): + results.append((name, bool(ok))) + print(("PASS " if ok else "FAIL ") + name + (f" [{detail}]" if detail else ""), flush=True) + + +# fake classifier: keep / ephemeral by marker, default otherwise +CLASSIFIED = [] + + +class Classifier(BaseHTTPRequestHandler): + def do_POST(self): + body = json.loads(self.rfile.read(int(self.headers["Content-Length"]))) + CLASSIFIED.append(body) + ex = body.get("excerpt", "") + cls = "keep" if "KEEPME" in ex else "ephemeral" if "EPHEMME" in ex else "default" + out = json.dumps({"class": cls}).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(out))) + self.end_headers() + self.wfile.write(out) + + def log_message(self, *a): + pass + + +threading.Thread(target=HTTPServer(("127.0.0.1", 8899), Classifier).serve_forever, daemon=True).start() + + +def start(tag, *extra): + global proc + a = [BIN, "-m", M, "-c", "90112", "-ngl", "99", "-fa", "on", "-np", "2", "-kvu", "-ctk", "q8_0", "-ctv", "q8_0", + "-ctkd", "q8_0", "-ctvd", "q8_0", "-ub", "1024", "-bs", "--jinja", "--spec-type", "draft-mtp", "--spec-draft-n-max", "2", + "--host", "127.0.0.1", "--port", "8898", "--cache-disk", DIR, *extra] + env = dict(os.environ, LD_LIBRARY_PATH="/run/opengl-driver/lib", GGML_CUDA_BATCH_INVARIANT="1") + proc = subprocess.Popen(a, stdout=open(f"{LOG}/{tag}.log", "w"), stderr=subprocess.STDOUT, env=env) + for _ in range(300): + try: + if urllib.request.urlopen(URL + "/health", timeout=2).status == 200: + return + except Exception: + pass + if proc.poll() is not None: + raise SystemExit(f"server exited, see {LOG}/{tag}.log") + time.sleep(1) + raise SystemExit("server did not start") + + +def stop(): + global proc + proc.terminate() + rc = proc.wait(600) + proc = None + return rc + + +def req(path, body=None, method=None): + r = urllib.request.Request(URL + path, data=None if body is None else json.dumps(body).encode(), + headers={"Content-Type": "application/json"}, method=method or ("POST" if body is not None else "GET")) + try: + with urllib.request.urlopen(r, timeout=3600) as f: + return f.status, json.loads(f.read()) + except urllib.error.HTTPError as e: + return e.code, json.loads(e.read() or b"{}") + + +def post(path): + return req(path, method="POST") + + +def text(marker, n_chars=5000, offset=0): + return f"{marker} " + FILLER[offset:offset + n_chars] + + +def chat(msgs, cache=None, ttl=None, n=16, cache_prompt=True): + if isinstance(msgs, str): + msgs = [{"role": "user", "content": msgs}] + body = {"messages": msgs, "max_tokens": n, "temperature": 0, "chat_template_kwargs": {"enable_thinking": False}, + "cache_prompt": cache_prompt} + if cache is not None: + body["cache"] = cache + if ttl is not None: + body["cache_ttl"] = ttl + st, x = req("/v1/chat/completions", body) + assert st == 200, x + return x["choices"][0]["message"]["content"], x["timings"]["prompt_n"] + + +def follow(msgs, ans): + if isinstance(msgs, str): + msgs = [{"role": "user", "content": msgs}] + return msgs + [{"role": "assistant", "content": ans}, {"role": "user", "content": "Now one sentence: what is the main point?"}] + + +def flush(): + # a skip request starts a task, which saves (and with -kvu clears) the idle slots + chat("FLUSH", cache="skip", n=1) + return wait_disk() + + +def wait_disk(): + for _ in range(600): + c = cache() + if not c.get("disk") or c["disk"]["pending"] == 0: + return c + time.sleep(0.5) + + +def cache(): + return req("/cache")[1] + + +def entries(c, tier, marker=""): + t = c.get(tier) or {} + return [e for e in t.get("entries", []) if e["excerpt"].startswith(marker)] + + +def markers(c, tier): + return sorted(e["excerpt"].split(" ")[0] for e in entries(c, tier) if e["excerpt"].startswith("EV-")) + + +def file_of(entry_id): + for p in glob.glob(DIR + "/*.kvc"): + with open(p, "rb") as f: + h = f.read(24) + if struct.unpack(" 0, e) + size = e[0]["bytes"] + keys = {"id", "class", "source", "classified", "n_tokens", "bytes", "hits", "age", "expires_in", "excerpt"} + check("GET /cache entry fields", keys <= set(e[0]) and {"by_class", "pinned_bytes", "entries"} <= set(c["ram"]) and "families" in c, sorted(e[0])) + + # skip: never stored, same greedy output + out_skip, _ = chat(text("CONV-skip"), cache="skip"); c = (flush(), cache())[1] + check("skip is never stored", not entries(c, "ram", "CONV-skip") and not entries(c, "disk", "CONV-skip")) + out_plain, _ = chat(text("CONV-skip")); c = (flush(), cache())[1] + check("greedy: skip == default", out_skip == out_plain, (out_skip, out_plain)) + x = entries(c, "ram", "CONV-skip") + check("default request stored", len(x) == 1) + xid = x[0]["id"] if x else None + + # a skip follow-up restores from the entry without consuming it + f = follow(text("CONV-skip"), out_plain) + out_f_skip, pn_skip = chat(f, cache="skip"); c = (flush(), cache())[1] + check("skip follow-up used the cache", pn_skip < 400, pn_skip) + check("skip leaves the entry in place", [e["id"] for e in entries(c, "ram", "CONV-skip")] == [xid]) + out_f_cold, pn_cold = chat(f, cache="skip", cache_prompt=False) + check("greedy: cached follow-up == cold follow-up", out_f_skip == out_f_cold, (out_f_skip, out_f_cold, pn_cold)) + out_f, pn = chat(f); c = (flush(), cache())[1] + check("greedy: normal follow-up == skip follow-up", out_f == out_f_skip and pn < 400, (out_f, pn)) + y = entries(c, "ram", "CONV-skip") + # the skip and cold follow-ups continued it too + check("continuation replaces the entry and counts its hits", len(y) == 1 and y[0]["id"] != xid and y[0]["hits"] >= 1, y) + + # ttl + chat(text("CONV-eph"), cache="ephemeral"); chat(text("CONV-ttl"), ttl=10); c = (flush(), cache())[1] + e1, e2 = entries(c, "ram", "CONV-eph"), entries(c, "ram", "CONV-ttl") + check("ephemeral stored with its ttl", len(e1) == 1 and e1[0]["class"] == "ephemeral" and 0 < e1[0]["expires_in"] <= 20, e1) + check("cache_ttl applies", len(e2) == 1 and 0 < e2[0]["expires_in"] <= 10, e2) + + # classifier + n0 = len(CLASSIFIED) + chat(text("CONV-cls KEEPME")); chat(text("CONV-cls2 EPHEMME")); flush() + for _ in range(60): + c = cache() + k1, k2 = entries(c, "ram", "CONV-cls "), entries(c, "ram", "CONV-cls2") + if k1 and k2 and k1[0]["source"] == "classifier" and k2[0]["source"] == "classifier": + break + time.sleep(1) + check("classifier applies classes", k1 and k1[0]["class"] == "keep" and k1[0]["classified"] and k2 and k2[0]["class"] == "ephemeral", (k1, k2)) + sent = [b for b in CLASSIFIED[n0:] if "KEEPME" in b["excerpt"]] + check("classifier request contract", sent and sent[0]["n_tokens"] > 1000 and sent[0]["n_turns"] >= 1 and len(sent[0]["excerpt"]) > 1000, + {k: (v if k != "excerpt" else len(v)) for k, v in (sent[0] if sent else {}).items()}) + + time.sleep(22) + c = cache() + check("expired entries dropped", not entries(c, "ram", "CONV-eph") and not entries(c, "ram", "CONV-ttl")) + + # family: one-shot calls behind one system prompt become ephemeral after 8 + sysmsg = {"role": "system", "content": "You are WEBX, a page summarizer. " + FILLER[200000:201500]} + for i in range(9): + chat([sysmsg, {"role": "user", "content": text(f"PAGE-{i}", 6000, 10000 * (i + 1))}], n=4) + c = (flush(), cache())[1] + fams = [x for x in c["families"] if x["n"] >= 8] + check("family learned ephemeral", fams and fams[0]["class"] == "ephemeral" and fams[0]["reuse_rate"] == 0, fams[:1]) + p8 = [e for e in c["ram"]["entries"] if e["source"] == "family"] + check("family class applied to entries", p8 and all(e["class"] == "ephemeral" for e in p8), len(p8)) + + # entry api + chat(text("CONV-pin")); chat(text("CONV-del")); c = (flush(), cache())[1] + pid, did = entries(c, "ram", "CONV-pin")[0]["id"], entries(c, "ram", "CONV-del")[0]["id"] + st, _ = post(f"/cache/entry?id={pid}&class=pin") + st2, _ = post(f"/cache/entry?id={did}&action=delete") + st3, _ = post("/cache/entry?id=0123456789abcdef&class=pin") + c = cache() + p = entries(c, "ram", "CONV-pin") + check("pin via api", st == 200 and p and p[0]["class"] == "pin" and p[0]["expires_in"] is None and c["ram"]["pinned_bytes"] > 0, p) + check("delete via api", st2 == 200 and not entries(c, "ram", "CONV-del")) + check("unknown id is 404", st3 == 404, st3) + + st, c = post("/cache/clear?scope=all") + check("clear keeps pinned", st == 200 and entries(c, "ram", "CONV-pin") and len(c["ram"]["entries"]) == 1, c.get("cleared")) + pinned = entries(c, "ram", "CONV-pin")[0] + + chat(text("CONV-v1")); flush() + print("stop run1", stop(), flush=True) + return size, pinned + + +def to_v1(src, dst, age): + with open(src, "rb") as f: + b = f.read() + with open(dst + ".part", "wb") as f: + f.write(b[:4] + struct.pack("