diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index d469635c..8755886b 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -2632,6 +2632,59 @@ private: res->n_erased = n_erased; queue_results.send(std::move(res)); } break; + case SERVER_TASK_TYPE_CACHE_GET: + case SERVER_TASK_TYPE_CACHE_CLEAR: + { + json cleared = json::object(); + if (task.type == SERVER_TASK_TYPE_CACHE_CLEAR) { + const std::string & scope = task.cache_scope; + const bool all = scope == "all"; + if (!all && scope != "slots" && scope != "ram" && scope != "disk") { + send_error(task, "scope must be slots, ram, disk or all", ERROR_TYPE_INVALID_REQUEST); + break; + } + if (all || scope == "slots") { + // a slot that is generating keeps its context; it is reported as busy + int n_cleared = 0; + int n_busy = 0; + for (auto & slot : slots) { + if (slot.is_processing()) { + n_busy++; + } else if (!slot.prompt.tokens.empty()) { + slot.prompt_clear(); + n_cleared++; + } + } + cleared["slots"] = n_cleared; + cleared["slots_busy"] = n_busy; + } + if (prompt_cache && (all || scope == "ram")) { + cleared["ram"] = prompt_cache->states.size(); + prompt_cache->states.clear(); + } + if (prompt_cache && (all || scope == "disk")) { + cleared["disk"] = prompt_cache->clear_disk(); + } + SRV_INF("cache cleared (%s): %s\n", scope.c_str(), cleared.dump().c_str()); + } + + json slots_info = json::array(); + for (const auto & slot : slots) { + slots_info.push_back({ + { "id", slot.id }, + { "n_tokens", slot.prompt.tokens.size() }, + { "processing", slot.is_processing() }, + }); + } + auto res = std::make_unique(); + res->id = task.id; + res->data = prompt_cache ? prompt_cache->info() : json::object(); + res->data["slots"] = slots_info; + if (task.type == SERVER_TASK_TYPE_CACHE_CLEAR) { + res->data["cleared"] = cleared; + } + queue_results.send(std::move(res)); + } break; case SERVER_TASK_TYPE_GET_LORA: { // TODO @ngxson : make lora_adapters a dedicated member of server_context @@ -4708,6 +4761,16 @@ void server_routes::init_routes() { return res; }; + // GET /cache: slots and both prompt-cache tiers; POST /cache/clear?scope=slots|ram|disk|all + this->get_cache = [this](const server_http_req & req) { + return handle_cache(req, SERVER_TASK_TYPE_CACHE_GET, ""); + }; + + this->post_cache_clear = [this](const server_http_req & req) { + const std::string scope = req.get_param("scope"); + return handle_cache(req, SERVER_TASK_TYPE_CACHE_CLEAR, scope.empty() ? "all" : scope); + }; + this->post_slots = [this](const server_http_req & req) { auto res = create_response(); if (params.slot_save_path.empty()) { @@ -5300,6 +5363,32 @@ std::unique_ptr server_routes::handle_slots_restore(const return res; } +std::unique_ptr server_routes::handle_cache(const server_http_req & req, server_task_type type, const std::string & scope) { + auto res = create_response(); + auto & rd = res->rd; + { + server_task task(type); + task.id = rd.get_new_id(); + task.cache_scope = scope; + rd.post_task(std::move(task), true); // high-priority task + } + + auto result = rd.next(req.should_stop); + if (!result) { + // connection was closed + GGML_ASSERT(req.should_stop()); + return res; + } + + if (result->is_error()) { + res->error(result->to_json()); + return res; + } + + res->ok(result->to_json()); + return res; +} + std::unique_ptr server_routes::handle_slots_erase(const server_http_req & req, int id_slot) { auto res = create_response(); auto & rd = res->rd; diff --git a/tools/server/server-context.h b/tools/server/server-context.h index 5d464b8e..adb1fa81 100644 --- a/tools/server/server-context.h +++ b/tools/server/server-context.h @@ -132,6 +132,8 @@ struct server_routes { server_http_context::handler_t get_metrics; server_http_context::handler_t get_slots; 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 get_props; server_http_context::handler_t post_props; server_http_context::handler_t post_infill; @@ -168,6 +170,7 @@ private: std::unique_ptr handle_slots_save(const server_http_req & req, int id_slot); std::unique_ptr handle_slots_restore(const server_http_req & req, int id_slot); std::unique_ptr handle_slots_erase(const server_http_req &, int id_slot); + std::unique_ptr handle_cache(const server_http_req & req, server_task_type type, const std::string & scope); std::unique_ptr handle_embeddings_impl(const server_http_req & req, task_response_type res_type); std::unique_ptr handle_count_tokens(const llama_vocab * vocab, mtmd_context * mctx, const server_http_req & req, task_response_type res_type); diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp index f858a541..ea3b0c81 100644 --- a/tools/server/server-task.cpp +++ b/tools/server/server-task.cpp @@ -2238,6 +2238,51 @@ void server_prompt_cache::disk_write(server_prompt_cache_state & state) { tokens.size(), 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 res = { + { "ram", { + { "prompts", states.size() }, + { "tokens", n_tokens() }, + { "bytes", size() }, + { "limit_bytes", limit_size }, + } }, + }; + if (disk_thread.joinable()) { + std::lock_guard lock(disk_mutex); + size_t bytes = 0; + size_t tokens = 0; + for (const auto & e : disk) { + bytes += e.size; + tokens += e.tokens.size(); + } + res["disk"] = { + { "dir", disk_dir }, + { "prompts", disk.size() }, + { "tokens", tokens }, + { "bytes", bytes }, + { "limit_bytes", disk_limit }, + { "pending", pending.size() }, + }; + } + return res; +} + +size_t server_prompt_cache::clear_disk() { + if (!disk_thread.joinable()) { + return 0; + } + std::unique_lock lock(disk_mutex); + disk_cv.wait(lock, [this] { return pending.empty(); }); + std::error_code ec; + for (const auto & e : disk) { + std::filesystem::remove(e.path, ec); + } + const size_t n = disk.size(); + disk.clear(); + SRV_INF("prompt cache: cleared %zu prompts from disk\n", n); + return n; +} + bool server_prompt_cache::disk_read(const std::string & path, bool has_mtmd, server_prompt_cache_state & out) const { disk_file f(path, "rb"); if (!f.f) { diff --git a/tools/server/server-task.h b/tools/server/server-task.h index b8d0b9ca..eb69e521 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -28,6 +28,8 @@ enum server_task_type { SERVER_TASK_TYPE_SLOT_SAVE, SERVER_TASK_TYPE_SLOT_RESTORE, SERVER_TASK_TYPE_SLOT_ERASE, + SERVER_TASK_TYPE_CACHE_GET, + SERVER_TASK_TYPE_CACHE_CLEAR, SERVER_TASK_TYPE_GET_LORA, SERVER_TASK_TYPE_SET_LORA, }; @@ -175,6 +177,9 @@ struct server_task { // used by SERVER_TASK_TYPE_METRICS bool metrics_reset_bucket = false; + // used by SERVER_TASK_TYPE_CACHE_CLEAR: "slots", "ram", "disk" or "all" + std::string cache_scope; + // used by SERVER_TASK_TYPE_SET_LORA std::map set_lora; // mapping adapter ID -> scale @@ -538,6 +543,15 @@ struct server_task_result_slot_erase : server_task_result { virtual json to_json() override; }; +// GET /cache and POST /cache/clear: the JSON is built on the server loop, which owns the slots and the prompt cache +struct server_task_result_cache : server_task_result { + json data; + + virtual json to_json() override { + return data; + } +}; + struct server_task_result_control : server_task_result { bool success = false; std::string message; // optional detail when success is false @@ -645,6 +659,12 @@ struct server_prompt_cache { void flush(); + // sizes and limits of both tiers + json info() const; + + // drops every disk entry and its file, after the writer has finished what is queued; returns the count + size_t clear_disk(); + private: struct disk_entry { std::string path; @@ -660,7 +680,7 @@ private: std::list disk; std::list pending; - std::mutex disk_mutex; + mutable std::mutex disk_mutex; std::condition_variable disk_cv; std::thread disk_thread; bool disk_stop = false; diff --git a/tools/server/server.cpp b/tools/server/server.cpp index 5fe2729b..3bb19082 100644 --- a/tools/server/server.cpp +++ b/tools/server/server.cpp @@ -219,6 +219,8 @@ int llama_server(common_params & params, int argc, char ** argv) { routes.post_lora_adapters = models_routes->proxy_post; routes.get_slots = models_routes->proxy_get; routes.post_slots = models_routes->proxy_post; + routes.get_cache = models_routes->proxy_get; + routes.post_cache_clear = models_routes->proxy_post; // custom routes for router routes.get_props = models_routes->get_router_props; @@ -272,6 +274,8 @@ int llama_server(common_params & params, int argc, char ** argv) { // Save & load slots ctx_http.get ("/slots", ex_wrapper(routes.get_slots)); 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)); // resumable streaming: a child binds the local session factories, the router binds // proxies that resolve the owning child, see server-stream.h