From fdb1db877c526ec90f668eca1b858da5dba85560 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adrien=20Gallou=C3=ABt?= Date: Thu, 2 Jul 2026 17:26:47 +0200 Subject: [PATCH 01/14] llama : add llama_model_ftype_name() (#25134) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * llama : add llama_model_ftype_name() Expose the model file type (quantization) name, e.g. "Q8_0" or "Q4_K - Medium", through a new public C API. The returned pointer is valid for the lifetime of the model and nullptr when the model is invalid or the file type is unknown. Signed-off-by: Adrien Gallouët * Export enum Signed-off-by: Adrien Gallouët * s/llama_model_ftype_name/llama_ftype_name/ Signed-off-by: Adrien Gallouët * Move "(guessed)" to the front in llama_ftype_name Prepend the "(guessed)" label instead of appending it. This allows removing the non-thread-safe static std::string, making the function allocation-free. Signed-off-by: Adrien Gallouët * Add LLAMA_FTYPE_PREFIX Signed-off-by: Adrien Gallouët * Dont check for model Signed-off-by: Adrien Gallouët --------- Signed-off-by: Adrien Gallouët --- include/llama.h | 6 +++ src/llama-model-loader.cpp | 90 +++++++++++++++++---------------- src/llama-model.cpp | 12 +++++ src/llama-model.h | 2 + tools/cli/cli.cpp | 3 ++ tools/server/server-context.cpp | 4 ++ tools/server/server-context.h | 1 + 7 files changed, 74 insertions(+), 44 deletions(-) diff --git a/include/llama.h b/include/llama.h index f723c9f60..b89ab758b 100644 --- a/include/llama.h +++ b/include/llama.h @@ -159,6 +159,9 @@ extern "C" { LLAMA_FTYPE_GUESSED = 1024, // not specified in the model file }; + // Get the model file type (quantization) as a string, e.g. "Q8_0" or "Q4_K - Medium" + LLAMA_API const char * llama_ftype_name(enum llama_ftype ftype); + enum llama_rope_scaling_type { LLAMA_ROPE_SCALING_TYPE_UNSPECIFIED = -1, LLAMA_ROPE_SCALING_TYPE_NONE = 0, @@ -606,6 +609,9 @@ extern "C" { // Get a string describing the model type LLAMA_API int32_t llama_model_desc(const struct llama_model * model, char * buf, size_t buf_size); + // Get the model file type (quantization), e.g. LLAMA_FTYPE_MOSTLY_Q8_0 + LLAMA_API enum llama_ftype llama_model_ftype(const struct llama_model * model); + // Returns the total size of all the tensors in the model in bytes LLAMA_API uint64_t llama_model_size(const struct llama_model * model); diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp index e07b0d231..55554735d 100644 --- a/src/llama-model-loader.cpp +++ b/src/llama-model-loader.cpp @@ -27,52 +27,54 @@ const char * llama_file_version_name(llama_fver version) { return "unknown"; } -static std::string llama_model_ftype_name(llama_ftype ftype) { - if (ftype & LLAMA_FTYPE_GUESSED) { - return llama_model_ftype_name((enum llama_ftype) (ftype & ~LLAMA_FTYPE_GUESSED)) + " (guessed)"; - } +#define LLAMA_FTYPE_PREFIX "(guessed) " - switch (ftype) { - case LLAMA_FTYPE_ALL_F32: return "all F32"; - case LLAMA_FTYPE_MOSTLY_F16: return "F16"; - case LLAMA_FTYPE_MOSTLY_BF16: return "BF16"; - case LLAMA_FTYPE_MOSTLY_Q1_0: return "Q1_0"; - case LLAMA_FTYPE_MOSTLY_Q4_0: return "Q4_0"; - case LLAMA_FTYPE_MOSTLY_Q4_1: return "Q4_1"; - case LLAMA_FTYPE_MOSTLY_Q5_0: return "Q5_0"; - case LLAMA_FTYPE_MOSTLY_Q5_1: return "Q5_1"; - case LLAMA_FTYPE_MOSTLY_Q8_0: return "Q8_0"; - case LLAMA_FTYPE_MOSTLY_MXFP4_MOE: return "MXFP4 MoE"; - case LLAMA_FTYPE_MOSTLY_NVFP4: return "NVFP4"; - case LLAMA_FTYPE_MOSTLY_Q2_K: return "Q2_K - Medium"; - case LLAMA_FTYPE_MOSTLY_Q2_K_S: return "Q2_K - Small"; - case LLAMA_FTYPE_MOSTLY_Q3_K_S: return "Q3_K - Small"; - case LLAMA_FTYPE_MOSTLY_Q3_K_M: return "Q3_K - Medium"; - case LLAMA_FTYPE_MOSTLY_Q3_K_L: return "Q3_K - Large"; - case LLAMA_FTYPE_MOSTLY_Q4_K_S: return "Q4_K - Small"; - case LLAMA_FTYPE_MOSTLY_Q4_K_M: return "Q4_K - Medium"; - case LLAMA_FTYPE_MOSTLY_Q5_K_S: return "Q5_K - Small"; - case LLAMA_FTYPE_MOSTLY_Q5_K_M: return "Q5_K - Medium"; - case LLAMA_FTYPE_MOSTLY_Q6_K: return "Q6_K"; - case LLAMA_FTYPE_MOSTLY_TQ1_0: return "TQ1_0 - 1.69 bpw ternary"; - case LLAMA_FTYPE_MOSTLY_TQ2_0: return "TQ2_0 - 2.06 bpw ternary"; - case LLAMA_FTYPE_MOSTLY_IQ2_XXS: return "IQ2_XXS - 2.0625 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ2_XS: return "IQ2_XS - 2.3125 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ2_S: return "IQ2_S - 2.5 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ2_M: return "IQ2_M - 2.7 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ3_XS: return "IQ3_XS - 3.3 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ3_XXS: return "IQ3_XXS - 3.0625 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ1_S: return "IQ1_S - 1.5625 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ1_M: return "IQ1_M - 1.75 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ4_NL: return "IQ4_NL - 4.5 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ4_XS: return "IQ4_XS - 4.25 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ3_S: return "IQ3_S - 3.4375 bpw"; - case LLAMA_FTYPE_MOSTLY_IQ3_M: return "IQ3_S mix - 3.66 bpw"; - - default: return "unknown, may not work"; +const char * llama_ftype_name(llama_ftype ftype) { + static constexpr size_t guessed_prefix_len = sizeof(LLAMA_FTYPE_PREFIX) - 1; + const char * name; + switch ((enum llama_ftype) (ftype & ~LLAMA_FTYPE_GUESSED)) { + case LLAMA_FTYPE_ALL_F32: name = LLAMA_FTYPE_PREFIX "all F32"; break; + case LLAMA_FTYPE_MOSTLY_F16: name = LLAMA_FTYPE_PREFIX "F16"; break; + case LLAMA_FTYPE_MOSTLY_BF16: name = LLAMA_FTYPE_PREFIX "BF16"; break; + case LLAMA_FTYPE_MOSTLY_Q1_0: name = LLAMA_FTYPE_PREFIX "Q1_0"; break; + case LLAMA_FTYPE_MOSTLY_Q4_0: name = LLAMA_FTYPE_PREFIX "Q4_0"; break; + case LLAMA_FTYPE_MOSTLY_Q4_1: name = LLAMA_FTYPE_PREFIX "Q4_1"; break; + case LLAMA_FTYPE_MOSTLY_Q5_0: name = LLAMA_FTYPE_PREFIX "Q5_0"; break; + case LLAMA_FTYPE_MOSTLY_Q5_1: name = LLAMA_FTYPE_PREFIX "Q5_1"; break; + case LLAMA_FTYPE_MOSTLY_Q8_0: name = LLAMA_FTYPE_PREFIX "Q8_0"; break; + case LLAMA_FTYPE_MOSTLY_MXFP4_MOE: name = LLAMA_FTYPE_PREFIX "MXFP4 MoE"; break; + case LLAMA_FTYPE_MOSTLY_NVFP4: name = LLAMA_FTYPE_PREFIX "NVFP4"; break; + case LLAMA_FTYPE_MOSTLY_Q2_K: name = LLAMA_FTYPE_PREFIX "Q2_K - Medium"; break; + case LLAMA_FTYPE_MOSTLY_Q2_K_S: name = LLAMA_FTYPE_PREFIX "Q2_K - Small"; break; + case LLAMA_FTYPE_MOSTLY_Q3_K_S: name = LLAMA_FTYPE_PREFIX "Q3_K - Small"; break; + case LLAMA_FTYPE_MOSTLY_Q3_K_M: name = LLAMA_FTYPE_PREFIX "Q3_K - Medium"; break; + case LLAMA_FTYPE_MOSTLY_Q3_K_L: name = LLAMA_FTYPE_PREFIX "Q3_K - Large"; break; + case LLAMA_FTYPE_MOSTLY_Q4_K_S: name = LLAMA_FTYPE_PREFIX "Q4_K - Small"; break; + case LLAMA_FTYPE_MOSTLY_Q4_K_M: name = LLAMA_FTYPE_PREFIX "Q4_K - Medium"; break; + case LLAMA_FTYPE_MOSTLY_Q5_K_S: name = LLAMA_FTYPE_PREFIX "Q5_K - Small"; break; + case LLAMA_FTYPE_MOSTLY_Q5_K_M: name = LLAMA_FTYPE_PREFIX "Q5_K - Medium"; break; + case LLAMA_FTYPE_MOSTLY_Q6_K: name = LLAMA_FTYPE_PREFIX "Q6_K"; break; + case LLAMA_FTYPE_MOSTLY_TQ1_0: name = LLAMA_FTYPE_PREFIX "TQ1_0 - 1.69 bpw ternary"; break; + case LLAMA_FTYPE_MOSTLY_TQ2_0: name = LLAMA_FTYPE_PREFIX "TQ2_0 - 2.06 bpw ternary"; break; + case LLAMA_FTYPE_MOSTLY_IQ2_XXS: name = LLAMA_FTYPE_PREFIX "IQ2_XXS - 2.0625 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ2_XS: name = LLAMA_FTYPE_PREFIX "IQ2_XS - 2.3125 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ2_S: name = LLAMA_FTYPE_PREFIX "IQ2_S - 2.5 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ2_M: name = LLAMA_FTYPE_PREFIX "IQ2_M - 2.7 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ3_XS: name = LLAMA_FTYPE_PREFIX "IQ3_XS - 3.3 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ3_XXS: name = LLAMA_FTYPE_PREFIX "IQ3_XXS - 3.0625 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ1_S: name = LLAMA_FTYPE_PREFIX "IQ1_S - 1.5625 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ1_M: name = LLAMA_FTYPE_PREFIX "IQ1_M - 1.75 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ4_NL: name = LLAMA_FTYPE_PREFIX "IQ4_NL - 4.5 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ4_XS: name = LLAMA_FTYPE_PREFIX "IQ4_XS - 4.25 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ3_S: name = LLAMA_FTYPE_PREFIX "IQ3_S - 3.4375 bpw"; break; + case LLAMA_FTYPE_MOSTLY_IQ3_M: name = LLAMA_FTYPE_PREFIX "IQ3_S mix - 3.66 bpw"; break; + default: name = LLAMA_FTYPE_PREFIX "unknown, may not work"; break; } + return (ftype & LLAMA_FTYPE_GUESSED) ? name : name + guessed_prefix_len; } +#undef LLAMA_FTYPE_PREFIX + // return a list of splits for a given path // for example, given "-00002-of-00004.gguf", returns list of all 4 splits static std::vector llama_get_list_splits(const std::string & path, const int idx, const int n_split) { @@ -1693,12 +1695,12 @@ bool llama_model_loader::load_all_data( } std::string llama_model_loader::ftype_name() const { - return llama_model_ftype_name(ftype); + return llama_ftype_name(ftype); } void llama_model_loader::print_info() const { LLAMA_LOG_INFO("%s: file format = %s\n", __func__, llama_file_version_name(fver)); - LLAMA_LOG_INFO("%s: file type = %s\n", __func__, llama_model_ftype_name(ftype).c_str()); + LLAMA_LOG_INFO("%s: file type = %s\n", __func__, llama_ftype_name(ftype)); if (n_bytes < GiB) { LLAMA_LOG_INFO("%s: file size = %.2f MiB (%.2f BPW) \n", __func__, n_bytes/1024.0/1024.0, n_bytes*8.0/n_elements); } else { diff --git a/src/llama-model.cpp b/src/llama-model.cpp index d58ebac28..e07f6e986 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -987,6 +987,8 @@ struct llama_model::impl { std::string desc_str; + llama_ftype ftype = LLAMA_FTYPE_ALL_F32; + // model memory mapped files llama_mmaps mappings; @@ -1200,6 +1202,8 @@ void llama_model_base::load_hparams(llama_model_loader & ml) { pimpl->desc_str = arch_name() + " " + type_name() + " " + ml.ftype_name(); + pimpl->ftype = ml.ftype; + if (hparams.f_max_alibi_bias > 0.0f) { hparams.use_alibi = true; } @@ -1646,6 +1650,10 @@ std::string llama_model::desc() const { return pimpl->desc_str; } +llama_ftype llama_model::ftype() const { + return pimpl->ftype; +} + size_t llama_model::size() const { return pimpl->n_bytes; } @@ -2616,6 +2624,10 @@ int32_t llama_model_desc(const llama_model * model, char * buf, size_t buf_size) return snprintf(buf, buf_size, "%s", model->desc().c_str()); } +llama_ftype llama_model_ftype(const llama_model * model) { + return model->ftype(); +} + uint64_t llama_model_size(const llama_model * model) { return model->size(); } diff --git a/src/llama-model.h b/src/llama-model.h index 4800d2928..45b054ced 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -637,6 +637,8 @@ struct llama_model { std::string desc() const; + llama_ftype ftype() const; + size_t size() const; // file size size_t n_tensors() const; size_t n_devices() const; diff --git a/tools/cli/cli.cpp b/tools/cli/cli.cpp index 8b7b58693..d974a4019 100644 --- a/tools/cli/cli.cpp +++ b/tools/cli/cli.cpp @@ -448,6 +448,9 @@ int llama_cli(int argc, char ** argv) { console::log("%s\n", LLAMA_ASCII_LOGO); console::log("build : %s\n", inf.build_info.c_str()); console::log("model : %s\n", inf.model_name.c_str()); + if (!inf.model_ftype.empty()) { + console::log("ftype : %s\n", inf.model_ftype.c_str()); + } console::log("modalities : %s\n", modalities.c_str()); if (!params.system_prompt.empty()) { console::log("using custom system prompt\n"); diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 39aa20b32..20e93258f 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -3989,6 +3989,8 @@ server_context_meta server_context::get_meta() const { auto bos_token_str = bos_id != LLAMA_TOKEN_NULL ? common_token_to_piece(impl->ctx_tgt, bos_id, true) : ""; auto eos_token_str = eos_id != LLAMA_TOKEN_NULL ? common_token_to_piece(impl->ctx_tgt, eos_id, true) : ""; + const char * ftype_name = llama_ftype_name(llama_model_ftype(impl->model_tgt)); + return server_context_meta { /* build_info */ std::string(llama_build_info()), /* model_name */ impl->model_name, @@ -4023,6 +4025,7 @@ server_context_meta server_context::get_meta() const { /* model_n_embd_inp */ llama_model_n_embd(impl->model_tgt), /* model_n_params */ llama_model_n_params(impl->model_tgt), /* model_size */ llama_model_size(impl->model_tgt), + /* model_ftype */ ftype_name, }; } @@ -5118,6 +5121,7 @@ json server_routes::get_model_info() const { {"n_embd", meta->model_n_embd_inp}, {"n_params", meta->model_n_params}, {"size", meta->model_size}, + {"ftype", meta->model_ftype}, }}, }; } diff --git a/tools/server/server-context.h b/tools/server/server-context.h index 952f825f7..f9ab1132b 100644 --- a/tools/server/server-context.h +++ b/tools/server/server-context.h @@ -50,6 +50,7 @@ struct server_context_meta { int32_t model_n_embd_inp; uint64_t model_n_params; uint64_t model_size; + std::string model_ftype; }; enum server_state { From c8ae9a750c6c89dc928503a16175b7b3c2d224e4 Mon Sep 17 00:00:00 2001 From: "Alessandro de Oliveira Faria (A.K.A.CABELO)" Date: Fri, 3 Jul 2026 05:26:54 -0300 Subject: [PATCH 02/14] vendor : update cpp-httplib to 0.49.0 (#25218) --- scripts/sync_vendor.py | 2 +- vendor/cpp-httplib/httplib.cpp | 162 ++++++++++++++++++++++++++------- vendor/cpp-httplib/httplib.h | 68 +++++++++++--- 3 files changed, 184 insertions(+), 48 deletions(-) diff --git a/scripts/sync_vendor.py b/scripts/sync_vendor.py index f913b0c7d..f66e78d63 100755 --- a/scripts/sync_vendor.py +++ b/scripts/sync_vendor.py @@ -5,7 +5,7 @@ import os import sys import subprocess -HTTPLIB_VERSION = "refs/tags/v0.48.0" +HTTPLIB_VERSION = "refs/tags/v0.49.0" vendor = { "https://github.com/nlohmann/json/releases/latest/download/json.hpp": "vendor/nlohmann/json.hpp", diff --git a/vendor/cpp-httplib/httplib.cpp b/vendor/cpp-httplib/httplib.cpp index 1ac4fa4ba..d65b7921b 100644 --- a/vendor/cpp-httplib/httplib.cpp +++ b/vendor/cpp-httplib/httplib.cpp @@ -478,7 +478,7 @@ bool set_socket_opt_time(socket_t sock, int level, int optname, } bool is_hex(char c, int &v) { - if (isdigit(static_cast(c))) { + if (is_ascii_digit(c)) { v = c - '0'; return true; } else if ('A' <= c && c <= 'F') { @@ -695,7 +695,11 @@ std::string base64_encode(const std::string &in) { std::string out; out.reserve(in.size()); - auto val = 0; + // Unsigned: the accumulator is never masked, so with a signed int the + // `val << 8` below overflows once enough bytes are folded in (undefined + // behaviour before C++20). Only the low bits are ever emitted, so the + // wrap-around of an unsigned accumulator does not affect the output. + uint32_t val = 0; auto valb = -6; for (auto c : in) { @@ -3887,8 +3891,7 @@ bool parse_range_header(const std::string &s, Ranges &ranges) { bool parse_range_header(const std::string &s, Ranges &ranges) try { #endif auto is_valid = [](const std::string &str) { - return std::all_of(str.cbegin(), str.cend(), - [](unsigned char c) { return std::isdigit(c); }); + return std::all_of(str.cbegin(), str.cend(), is_ascii_digit); }; if (s.size() > 7 && s.compare(0, 6, "bytes=") == 0) { @@ -4336,7 +4339,7 @@ bool is_multipart_boundary_chars_valid(const std::string &boundary) { auto valid = true; for (size_t i = 0; i < boundary.size(); i++) { auto c = boundary[i]; - if (!std::isalnum(static_cast(c)) && c != '-' && c != '_') { + if (!is_ascii_alnum(c) && c != '-' && c != '_') { valid = false; break; } @@ -4344,18 +4347,47 @@ bool is_multipart_boundary_chars_valid(const std::string &boundary) { return valid; } +// Escape a multipart field name/filename following the WHATWG HTML standard +// ("escape a multipart form-data name"), which is what browsers send: +// '"' -> %22, CR -> %0D, LF -> %0A +// With escape_quote = false, only CR and LF are escaped; this is for header +// values outside a quoted-string (e.g. Content-Type), where '"' is legal. +std::string escape_multipart_field(const std::string &s, + bool escape_quote = true) { + std::string result; + result.reserve(s.size()); + for (auto c : s) { + switch (c) { + case '"': + if (escape_quote) { + result += "%22"; + } else { + result += c; + } + break; + case '\r': result += "%0D"; break; + case '\n': result += "%0A"; break; + default: result += c; break; + } + } + return result; +} + template std::string serialize_multipart_formdata_item_begin(const T &item, const std::string &boundary) { std::string body = "--" + boundary + "\r\n"; - body += "Content-Disposition: form-data; name=\"" + item.name + "\""; + body += "Content-Disposition: form-data; name=\"" + + escape_multipart_field(item.name) + "\""; if (!item.filename.empty()) { - body += "; filename=\"" + item.filename + "\""; + body += "; filename=\"" + escape_multipart_field(item.filename) + "\""; } body += "\r\n"; if (!item.content_type.empty()) { - body += "Content-Type: " + item.content_type + "\r\n"; + body += + "Content-Type: " + escape_multipart_field(item.content_type, false) + + "\r\n"; } body += "\r\n"; @@ -4821,10 +4853,9 @@ private: namespace fields { bool is_token_char(char c) { - return std::isalnum(static_cast(c)) || c == '!' || c == '#' || - c == '$' || c == '%' || c == '&' || c == '\'' || c == '*' || - c == '+' || c == '-' || c == '.' || c == '^' || c == '_' || c == '`' || - c == '|' || c == '~'; + return is_ascii_alnum(c) || c == '!' || c == '#' || c == '$' || c == '%' || + c == '&' || c == '\'' || c == '*' || c == '+' || c == '-' || + c == '.' || c == '^' || c == '_' || c == '`' || c == '|' || c == '~'; } bool is_token(const std::string &s) { @@ -4873,7 +4904,8 @@ bool is_field_value(const std::string &s) { return is_field_content(s); } } // namespace fields bool perform_websocket_handshake(Stream &strm, const std::string &host, - int port, const std::string &path, + int port, bool is_ssl, + const std::string &path, const Headers &headers, std::string &selected_subprotocol) { // Validate path and host @@ -4899,7 +4931,7 @@ bool perform_websocket_handshake(Stream &strm, const std::string &host, // Build upgrade request std::string req_str = "GET " + path + " HTTP/1.1\r\n"; - req_str += "Host: " + host + ":" + std::to_string(port) + "\r\n"; + req_str += "Host: " + make_host_and_port_string(host, port, is_ssl) + "\r\n"; req_str += "Upgrade: websocket\r\n"; req_str += "Connection: Upgrade\r\n"; req_str += "Sec-WebSocket-Key: " + client_key + "\r\n"; @@ -5599,9 +5631,8 @@ std::string encode_uri_component(const std::string &value) { escaped << std::hex; for (auto c : value) { - if (std::isalnum(static_cast(c)) || c == '-' || c == '_' || - c == '.' || c == '!' || c == '~' || c == '*' || c == '\'' || c == '(' || - c == ')') { + if (detail::is_ascii_alnum(c) || c == '-' || c == '_' || c == '.' || + c == '!' || c == '~' || c == '*' || c == '\'' || c == '(' || c == ')') { escaped << c; } else { escaped << std::uppercase; @@ -5620,10 +5651,10 @@ std::string encode_uri(const std::string &value) { escaped << std::hex; for (auto c : value) { - if (std::isalnum(static_cast(c)) || c == '-' || c == '_' || - c == '.' || c == '!' || c == '~' || c == '*' || c == '\'' || c == '(' || - c == ')' || c == ';' || c == '/' || c == '?' || c == ':' || c == '@' || - c == '&' || c == '=' || c == '+' || c == '$' || c == ',' || c == '#') { + if (detail::is_ascii_alnum(c) || c == '-' || c == '_' || c == '.' || + c == '!' || c == '~' || c == '*' || c == '\'' || c == '(' || c == ')' || + c == ';' || c == '/' || c == '?' || c == ':' || c == '@' || c == '&' || + c == '=' || c == '+' || c == '$' || c == ',' || c == '#') { escaped << c; } else { escaped << std::uppercase; @@ -5684,7 +5715,8 @@ std::string encode_path_component(const std::string &component) { auto c = static_cast(component[i]); // Unreserved characters per RFC 3986: ALPHA / DIGIT / "-" / "." / "_" / "~" - if (std::isalnum(c) || c == '-' || c == '.' || c == '_' || c == '~') { + if (detail::is_ascii_alnum(static_cast(c)) || c == '-' || c == '.' || + c == '_' || c == '~') { result += static_cast(c); } // Path-safe sub-delimiters: "!" / "$" / "&" / "'" / "(" / ")" / "*" / "+" / @@ -5757,7 +5789,8 @@ std::string encode_query_component(const std::string &component, auto c = static_cast(component[i]); // Unreserved characters per RFC 3986 - if (std::isalnum(c) || c == '-' || c == '.' || c == '_' || c == '~') { + if (detail::is_ascii_alnum(static_cast(c)) || c == '-' || c == '.' || + c == '_' || c == '~') { result += static_cast(c); } // Space handling @@ -6010,6 +6043,48 @@ size_t MultipartFormData::get_file_count(const std::string &key) const { return static_cast(std::distance(r.first, r.second)); } +// Multipart FormData writer implementation +bool is_valid_multipart_boundary(const std::string &boundary) { + return detail::is_multipart_boundary_chars_valid(boundary); +} + +MultipartFormDataWriter::MultipartFormDataWriter() + : boundary_(detail::make_multipart_data_boundary()) {} + +MultipartFormDataWriter::MultipartFormDataWriter(std::string boundary) + : boundary_(std::move(boundary)) {} + +const std::string &MultipartFormDataWriter::boundary() const { + return boundary_; +} + +std::string MultipartFormDataWriter::content_type() const { + return detail::serialize_multipart_formdata_get_content_type(boundary_); +} + +std::string +MultipartFormDataWriter::serialize(const UploadFormDataItems &items) const { + return detail::serialize_multipart_formdata(items, boundary_); +} + +size_t MultipartFormDataWriter::content_length( + const UploadFormDataItems &items) const { + return detail::get_multipart_content_length(items, boundary_); +} + +std::string +MultipartFormDataWriter::item_begin(const UploadFormData &item) const { + return detail::serialize_multipart_formdata_item_begin(item, boundary_); +} + +std::string MultipartFormDataWriter::item_end() { + return detail::serialize_multipart_formdata_item_end(); +} + +std::string MultipartFormDataWriter::finish() const { + return detail::serialize_multipart_formdata_finish(boundary_); +} + // Response implementation size_t Response::get_header_value_u64(const std::string &key, size_t def, size_t id) const { @@ -6229,8 +6304,10 @@ ssize_t detail::BodyReader::read(char *buf, size_t len) { } // ThreadPool implementation -ThreadPool::ThreadPool(size_t n, size_t max_n, size_t mqr) - : base_thread_count_(n), max_queued_requests_(mqr), idle_thread_count_(0), +ThreadPool::ThreadPool(size_t n, size_t max_n, size_t mqr, + time_t idle_timeout_sec) + : base_thread_count_(n), max_queued_requests_(mqr), + idle_timeout_sec_(idle_timeout_sec), idle_thread_count_(0), shutdown_(false) { #ifndef CPPHTTPLIB_NO_EXCEPTIONS if (max_n != 0 && max_n < n) { @@ -6340,9 +6417,9 @@ void ThreadPool::worker(bool is_dynamic) { idle_thread_count_++; if (is_dynamic) { - auto has_work = cond_.wait_for( - lock, std::chrono::seconds(CPPHTTPLIB_THREAD_POOL_IDLE_TIMEOUT), - [&] { return !jobs_.empty() || shutdown_; }); + auto has_work = + cond_.wait_for(lock, std::chrono::seconds(idle_timeout_sec_), + [&] { return !jobs_.empty() || shutdown_; }); if (!has_work) { // Timed out with no work - exit this dynamic thread idle_thread_count_--; @@ -9687,9 +9764,18 @@ bool ClientImpl::write_request(Stream &strm, Request &req, if (!query_part.empty()) { // Normalize the query string (decode then re-encode) while preserving - // the original parameter order. - auto normalized = detail::normalize_query_string(query_part); - if (!normalized.empty()) { path_with_query += '?' + normalized; } + // the original parameter order. When path encoding is disabled the + // caller has supplied an already-encoded target and expects the exact + // bytes to be sent on the wire, so skip normalization for the query + // too. Normalizing here would decode-then-re-encode the query and + // corrupt pre-encoded binary payloads (e.g. turning `%20` into `+`, + // which a strict RFC 3986 server decodes back as `+`, not a space). + if (path_encode_) { + auto normalized = detail::normalize_query_string(query_part); + if (!normalized.empty()) { path_with_query += '?' + normalized; } + } else { + path_with_query += '?' + query_part; + } // Still populate req.params for handlers/users who read them. detail::parse_query_text(query_part, req.params); @@ -12518,7 +12604,7 @@ bool is_ipv4_address(const std::string &str) { for (char c : str) { if (c == '.') { dots++; - } else if (!isdigit(static_cast(c))) { + } else if (!detail::is_ascii_digit(c)) { return false; } } @@ -12535,7 +12621,7 @@ bool parse_ipv4(const std::string &str, unsigned char *out) { } int val = 0; int digits = 0; - while (*p >= '0' && *p <= '9') { + while (detail::is_ascii_digit(*p)) { val = val * 10 + (*p - '0'); if (val > 255) { return false; } p++; @@ -16487,9 +16573,15 @@ bool WebSocketClient::connect() { return false; } +#ifdef CPPHTTPLIB_SSL_ENABLED + auto is_ssl = is_ssl_; +#else + auto is_ssl = false; +#endif + std::string selected_subprotocol; - if (!detail::perform_websocket_handshake(*strm, host_, port_, path_, headers_, - selected_subprotocol)) { + if (!detail::perform_websocket_handshake(*strm, host_, port_, is_ssl, path_, + headers_, selected_subprotocol)) { shutdown_and_close(); return false; } diff --git a/vendor/cpp-httplib/httplib.h b/vendor/cpp-httplib/httplib.h index bfdbfc1da..e7ef56370 100644 --- a/vendor/cpp-httplib/httplib.h +++ b/vendor/cpp-httplib/httplib.h @@ -8,8 +8,8 @@ #ifndef CPPHTTPLIB_HTTPLIB_H #define CPPHTTPLIB_HTTPLIB_H -#define CPPHTTPLIB_VERSION "0.48.0" -#define CPPHTTPLIB_VERSION_NUM "0x003000" +#define CPPHTTPLIB_VERSION "0.49.0" +#define CPPHTTPLIB_VERSION_NUM "0x003100" #ifdef _WIN32 #if defined(_WIN32_WINNT) && _WIN32_WINNT < 0x0A00 @@ -309,7 +309,6 @@ using socket_t = int; #include #include #include -#include #include #include #include @@ -540,6 +539,21 @@ make_unique(std::size_t n) { return std::unique_ptr(new RT[n]); } +// Locale-independent ASCII character classification. The +// counterparts (std::isalnum, std::isdigit, ...) consult the global C locale, +// so e.g. std::isalnum(0xC5) can return true once an embedder calls +// setlocale(). HTTP grammars are defined over ASCII, so raw bytes must be +// classified without regard to the locale. +inline bool is_ascii_digit(char c) { return '0' <= c && c <= '9'; } + +inline bool is_ascii_alpha(char c) { + return ('a' <= c && c <= 'z') || ('A' <= c && c <= 'Z'); +} + +inline bool is_ascii_alnum(char c) { + return is_ascii_digit(c) || is_ascii_alpha(c); +} + namespace case_ignore { inline unsigned char to_lower(int c) { @@ -661,7 +675,7 @@ inline from_chars_result from_chars(const char *first, const char *last, for (; p != last; ++p) { char c = *p; int digit = -1; - if ('0' <= c && c <= '9') { + if (is_ascii_digit(c)) { digit = c - '0'; } else if ('a' <= c && c <= 'z') { digit = c - 'a' + 10; @@ -733,14 +747,14 @@ inline from_chars_result from_chars(const char *first, const char *last, return false; }; - for (; p != last && '0' <= *p && *p <= '9'; ++p) { + for (; p != last && is_ascii_digit(*p); ++p) { seen_digit = true; accumulate(*p); } if (p != last && *p == '.') { ++p; - for (; p != last && '0' <= *p && *p <= '9'; ++p) { + for (; p != last && is_ascii_digit(*p); ++p) { seen_digit = true; if (frac_digits < max_frac_digits && accumulate(*p)) { ++frac_digits; } } @@ -803,8 +817,8 @@ inline bool parse_url(const std::string &url, UrlComponents &uc) { // IPv6 host must be [a-fA-F0-9:]+ only if (uc.host.empty()) { return false; } for (auto c : uc.host) { - if (!((c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F') || - (c >= '0' && c <= '9') || c == ':')) { + if (!(is_ascii_digit(c) || (c >= 'a' && c <= 'f') || + (c >= 'A' && c <= 'F') || c == ':')) { return false; } } @@ -1541,7 +1555,9 @@ public: class ThreadPool final : public TaskQueue { public: - explicit ThreadPool(size_t n, size_t max_n = 0, size_t mqr = 0); + explicit ThreadPool( + size_t n, size_t max_n = 0, size_t mqr = 0, + time_t idle_timeout_sec = CPPHTTPLIB_THREAD_POOL_IDLE_TIMEOUT); ThreadPool(const ThreadPool &) = delete; ~ThreadPool() override = default; @@ -1556,6 +1572,7 @@ private: size_t base_thread_count_; size_t max_thread_count_; size_t max_queued_requests_; + time_t idle_timeout_sec_; size_t idle_thread_count_; bool shutdown_; @@ -1680,6 +1697,35 @@ make_multipart_content_provider(const UploadFormDataItems &items, } // namespace detail +bool is_valid_multipart_boundary(const std::string &boundary); + +// Serializer for multipart/form-data request bodies. The boundary is owned +// by the writer so that per-part framing and the final terminator always +// agree. Field names and filenames are escaped following the WHATWG HTML +// standard ('"' -> %22, CR -> %0D, LF -> %0A); CR and LF are also escaped +// in content types. +class MultipartFormDataWriter { +public: + MultipartFormDataWriter(); + // precondition: is_valid_multipart_boundary(boundary) + explicit MultipartFormDataWriter(std::string boundary); + + const std::string &boundary() const; + std::string content_type() const; + + // In-memory items -> whole body (known length) + std::string serialize(const UploadFormDataItems &items) const; + size_t content_length(const UploadFormDataItems &items) const; + + // Per-part framing for streaming via a content provider + std::string item_begin(const UploadFormData &item) const; + static std::string item_end(); + std::string finish() const; + +private: + std::string boundary_; +}; + class Server { public: using Handler = std::function; @@ -2897,9 +2943,7 @@ template inline constexpr size_t str_len(const char (&)[N]) { } inline bool is_numeric(const std::string &str) { - return !str.empty() && - std::all_of(str.cbegin(), str.cend(), - [](unsigned char c) { return std::isdigit(c); }); + return !str.empty() && std::all_of(str.cbegin(), str.cend(), is_ascii_digit); } inline size_t get_header_value_u64(const Headers &headers, From 5a460dea9f961cdb508d58a6e7b0f9e259b4c19f Mon Sep 17 00:00:00 2001 From: Gaurav Garg Date: Fri, 3 Jul 2026 14:36:29 +0530 Subject: [PATCH 03/14] Remove redundant CUDA copies after gated_delta_net. (#23940) * Remove redundant CUDA copies after gated_delta_net. Currently, GDN writes recurrent state snapshots into its output tail, then the graph immediately copies those snapshots into ssm_states_all. With MTP draft length 3, target decode uses K=4, so that becomes 4 extra ggml_cuda_cpy calls. The change detects that gated_delta_net -> view -> cpy pattern and makes the CUDA GDN kernel write the state snapshot(s) directly into the recurrent cache, skipping the intermediate tail writes and copy kernels when safe. * Address review comments --- ggml/src/ggml-cuda/gated_delta_net.cu | 65 +++++++++++-------- ggml/src/ggml-cuda/gated_delta_net.cuh | 10 +++ ggml/src/ggml-cuda/ggml-cuda.cu | 87 +++++++++++++++++++++++++- 3 files changed, 135 insertions(+), 27 deletions(-) diff --git a/ggml/src/ggml-cuda/gated_delta_net.cu b/ggml/src/ggml-cuda/gated_delta_net.cu index a547360eb..1b431a724 100644 --- a/ggml/src/ggml-cuda/gated_delta_net.cu +++ b/ggml/src/ggml-cuda/gated_delta_net.cu @@ -10,6 +10,7 @@ gated_delta_net_cuda(const float * q, const float * beta, const float * curr_state, float * dst, + float * state, int64_t H, int64_t n_tokens, int64_t n_seqs, @@ -25,6 +26,7 @@ gated_delta_net_cuda(const float * q, const uint3 neqk1_magic, const uint3 rq3_magic, float scale, + int64_t state_slot_stride, int K) { const uint32_t h_idx = blockIdx.x; const uint32_t sequence = blockIdx.y; @@ -35,9 +37,7 @@ gated_delta_net_cuda(const float * q, const uint32_t iq1 = fastmodulo(h_idx, neqk1_magic); const uint32_t iq3 = fastdiv(sequence, rq3_magic); - const int64_t attn_score_elems = S_v * H * n_tokens * n_seqs; float * attn_data = dst; - float * state = dst + attn_score_elems; // input state holds s0 only: [S_v, S_v, H, n_seqs] — seq stride is D = H * S_v * S_v. // output state layout (per-slot D * n_seqs) — same per-(seq,head) offset as before. @@ -145,10 +145,9 @@ gated_delta_net_cuda(const float * q, if constexpr (keep_rs_t) { // snapshot slot mapping: slot 0 = most recent state, slot s = s tokens back. // When n_tokens < K only slots 0..n_tokens-1 are written; older slots are caller-owned. - const int64_t state_size_per_token = S_v * S_v * H * n_seqs; // per-slot stride in output const int target_slot = (int) n_tokens - 1 - t; if (target_slot >= 0 && target_slot < K) { - float * curr_state = (dst + attn_score_elems) + target_slot * state_size_per_token + state_out_offset; + float * curr_state = state + target_slot * state_slot_stride; #pragma unroll for (int r = 0; r < rows_per_lane; r++) { const int i = r * warp_size + lane; @@ -171,13 +170,13 @@ template static void launch_gated_delta_net( const float * q_d, const float * k_d, const float * v_d, const float * g_d, const float * b_d, const float * s_d, - float * dst_d, + float * dst_d, float * state_d, int64_t S_v, int64_t H, int64_t n_tokens, int64_t n_seqs, int64_t sq1, int64_t sq2, int64_t sq3, int64_t sv1, int64_t sv2, int64_t sv3, int64_t sb1, int64_t sb2, int64_t sb3, int64_t neqk1, int64_t rq3, - float scale, int K, cudaStream_t stream) { + float scale, int64_t state_slot_stride, int K, cudaStream_t stream) { //TODO: Add chunked kernel for even faster pre-fill const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size; const int num_warps = 4; @@ -187,34 +186,32 @@ static void launch_gated_delta_net( const uint3 neqk1_magic = init_fastdiv_values(neqk1); const uint3 rq3_magic = init_fastdiv_values(rq3); - int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; - const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, stream); switch (S_v) { case 16: ggml_cuda_kernel_launch(gated_delta_net_cuda<16, KDA, keep_rs_t>, launch_params, - q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, + q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); break; case 32: ggml_cuda_kernel_launch(gated_delta_net_cuda<32, KDA, keep_rs_t>, launch_params, - q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, + q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); break; case 64: { ggml_cuda_kernel_launch(gated_delta_net_cuda<64, KDA, keep_rs_t>, launch_params, - q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, + q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); break; } case 128: { ggml_cuda_kernel_launch(gated_delta_net_cuda<128, KDA, keep_rs_t>, launch_params, - q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, + q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); break; } default: @@ -223,7 +220,8 @@ static void launch_gated_delta_net( } } -void ggml_cuda_op_gated_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { +static void ggml_cuda_op_gated_delta_net_impl( + ggml_backend_cuda_context & ctx, ggml_tensor * dst, const ggml_cuda_gated_delta_net_fused_cache * cache) { ggml_tensor * src_q = dst->src[0]; ggml_tensor * src_k = dst->src[1]; ggml_tensor * src_v = dst->src[2]; @@ -288,25 +286,42 @@ void ggml_cuda_op_gated_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * const int K = ggml_get_op_params_i32(dst, 0); const bool keep_rs = K > 1; + // recurrent state -> gdn_out tail (after attention scores), or the cache when fusing + float * state_d = dst_d + S_v * H * n_tokens * n_seqs; + int64_t state_slot_stride = S_v * S_v * H * n_seqs; + if (cache != nullptr) { + state_d = cache->data; + state_slot_stride = cache->slot_stride; + } + if (kda) { if (keep_rs) { - launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, + launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, S_v, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1, rq3, scale, K, stream); + sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, stream); } else { - launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, + launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, S_v, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1, rq3, scale, K, stream); + sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, stream); } } else { if (keep_rs) { - launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, + launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, S_v, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1, rq3, scale, K, stream); + sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, stream); } else { - launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, + launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, S_v, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1, rq3, scale, K, stream); + sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, stream); } } } + +void ggml_cuda_op_gated_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + ggml_cuda_op_gated_delta_net_impl(ctx, dst, nullptr); +} + +void ggml_cuda_op_gated_delta_net_fused_cache( + ggml_backend_cuda_context & ctx, ggml_tensor * dst, ggml_cuda_gated_delta_net_fused_cache cache) { + ggml_cuda_op_gated_delta_net_impl(ctx, dst, &cache); +} diff --git a/ggml/src/ggml-cuda/gated_delta_net.cuh b/ggml/src/ggml-cuda/gated_delta_net.cuh index 7375e81c0..f9bf43706 100644 --- a/ggml/src/ggml-cuda/gated_delta_net.cuh +++ b/ggml/src/ggml-cuda/gated_delta_net.cuh @@ -1,4 +1,14 @@ #include "common.cuh" #include "ggml.h" +// fused-kernel recurrent-state output; strides in elements (per-seq stride is always D, set in-kernel) +struct ggml_cuda_gated_delta_net_fused_cache { + float * data; // rollback slot 0 + int64_t slot_stride; // between rollback slots (0 when K==1) +}; + void ggml_cuda_op_gated_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst); + +// same op, but writes the snapshot(s) into the cache instead of dst (see ggml_cuda_try_gdn_cache_fusion) +void ggml_cuda_op_gated_delta_net_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst, + ggml_cuda_gated_delta_net_fused_cache cache); diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index cca70592f..78d2218e5 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -3251,6 +3251,11 @@ static void ggml_backend_cuda_synchronize(ggml_backend_t backend) { GGML_UNUSED(backend); } +static bool ggml_cuda_is_view_or_noop(const ggml_tensor * t) { + return ggml_is_empty(t) || t->op == GGML_OP_RESHAPE || t->op == GGML_OP_TRANSPOSE || + t->op == GGML_OP_VIEW || t->op == GGML_OP_PERMUTE || t->op == GGML_OP_NONE; +} + #ifdef USE_CUDA_GRAPH static bool ggml_cuda_graph_check_compability(ggml_cgraph * cgraph) { @@ -3260,7 +3265,7 @@ static bool ggml_cuda_graph_check_compability(ggml_cgraph * cgraph) { for (int i = 0; i < cgraph->n_nodes; i++) { ggml_tensor * node = cgraph->nodes[i]; - if (ggml_is_empty(node) || node->op == GGML_OP_RESHAPE || node->op == GGML_OP_TRANSPOSE || node->op == GGML_OP_VIEW || node->op == GGML_OP_PERMUTE || node->op == GGML_OP_NONE) { + if (ggml_cuda_is_view_or_noop(node)) { continue; } @@ -3403,6 +3408,70 @@ static bool ggml_cuda_should_fuse_rope_set_rows(const ggml_tensor * rope, return true; } +// match gated_delta_net + the strided cpy that scatters its state snapshots into the cache +// (slot i -> rollback group i, slot 0 newest), so the kernel can write them and skip the cpy. +static int ggml_cuda_try_gdn_cache_fusion( + const ggml_cgraph * cgraph, int node_idx, ggml_cuda_gated_delta_net_fused_cache & fused_state_cpy) { + const ggml_tensor * gdn = cgraph->nodes[node_idx]; + // the kernel skips the snapshot tail, so the gdn output must not be a graph output + if (gdn->op != GGML_OP_GATED_DELTA_NET || gdn->type != GGML_TYPE_F32 || + (gdn->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return 0; + } + + const ggml_tensor * src_v = gdn->src[2]; + const int64_t S_v = src_v->ne[0]; + const int64_t H = src_v->ne[1]; + const int64_t n_tokens = src_v->ne[2]; + const int64_t n_seqs = src_v->ne[3]; + const int64_t D = S_v * S_v * H; + const int64_t K = ggml_get_op_params_i32(gdn, 0); // snapshot slot count + const int64_t n_written = std::min(n_tokens, K); // newest n_written slots are written + + // snapshot tail starts right after the attention scores + const size_t tail_off = ggml_row_size(GGML_TYPE_F32, S_v * H * n_tokens * n_seqs); + + // snapshot cpy is the first real node after the gdn (skip views/no-ops) + const ggml_tensor * cpy = nullptr; + int skip = 0; + for (int j = node_idx + 1; j < cgraph->n_nodes && cpy == nullptr; ++j) { + const ggml_tensor * n = cgraph->nodes[j]; + if (ggml_cuda_is_view_or_noop(n)) { + continue; + } + if (n->op != GGML_OP_CPY || (n->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return 0; + } + cpy = n; + skip = j - node_idx; + } + if (cpy == nullptr) { + return 0; + } + + const ggml_tensor * src = cpy->src[0]; // view of the gdn snapshot tail + const ggml_tensor * dst = cpy->src[1]; // cache view the kernel writes to + + // src must be this gdn's snapshot tail (contiguous, at the tail offset) + if (src->op != GGML_OP_VIEW || src->view_src != gdn || src->view_offs != tail_off || + !ggml_is_contiguous(src)) { + return 0; + } + + // dst is the [D, n_seqs, n_written] cache view; require nb[1] == D (the per-seq stride the kernel + // assumes). ggml_cpy pins src to the same element count. + const std::array expected_ne = { D, n_seqs, n_written, 1 }; + if (dst->op != GGML_OP_VIEW || dst->type != GGML_TYPE_F32 || dst->data == nullptr || + !std::equal(expected_ne.begin(), expected_ne.end(), dst->ne) || + dst->nb[0] != ggml_type_size(GGML_TYPE_F32) || dst->nb[1] != (size_t) ggml_row_size(GGML_TYPE_F32, D)) { + return 0; + } + + fused_state_cpy.data = (float *) dst->data; // rollback group 0 (newest) + fused_state_cpy.slot_stride = K > 1 ? (int64_t) (dst->nb[2] / sizeof(float)) : 0; + return skip; +} + static bool ggml_cuda_topk_moe_fusion(const struct ggml_cgraph * cgraph, int node_idx, ggml_cuda_topk_moe_args & args) { args.sigmoid = false; args.softmax = false; @@ -3844,6 +3913,20 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph ggml_tensor * node = cgraph->nodes[i]; + // gated_delta_net -> cpy: scatter recurrent-state snapshots into the cache + if (node->op == GGML_OP_GATED_DELTA_NET) { + ggml_cuda_gated_delta_net_fused_cache fused_state_cpy; + const int nodes_to_skip = ggml_cuda_try_gdn_cache_fusion(cgraph, i, fused_state_cpy); + if (nodes_to_skip > 0) { +#ifdef GGML_CUDA_DEBUG + GGML_LOG_INFO("%s: fused gated_delta_net snapshot copies for %s (skipped %d nodes)\n", + __func__, node->name, nodes_to_skip); +#endif + ggml_cuda_op_gated_delta_net_fused_cache(*cuda_ctx, node, fused_state_cpy); + return nodes_to_skip; + } + } + //topk-moe if (cgraph->nodes[i]->op == GGML_OP_UNARY || cgraph->nodes[i]->op == GGML_OP_SOFT_MAX || cgraph->nodes[i]->op == GGML_OP_ARGSORT) { @@ -4372,7 +4455,7 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud #endif prev_i = i; - if (ggml_is_empty(node) || node->op == GGML_OP_RESHAPE || node->op == GGML_OP_TRANSPOSE || node->op == GGML_OP_VIEW || node->op == GGML_OP_PERMUTE || node->op == GGML_OP_NONE) { + if (ggml_cuda_is_view_or_noop(node)) { continue; } From 94875285e47516da4e1c71e7bca5ba65b82bbe59 Mon Sep 17 00:00:00 2001 From: Aleksander Grygier Date: Fri, 3 Jul 2026 12:16:29 +0200 Subject: [PATCH 04/14] ui: Add MCP Servers Opt-In for first time visitors (#25239) * feat: ui: Add predefined recommended MCP servers to settings * feat: ui: Add MCP server recommendation dialog with custom server support * feat: Auto-focus input fields on mount and dynamic addition * feat: Add header validation to MCP server add and edit forms * feat: Persist recommended MCP server opt-in selections * test: Cover MCP configuration with tests * chore: Format & cleanup * feat: Centralize MCP server overrides to settings config and improve recommendation UI * fix: Capture index before mutation to prevent focus drift * refactor: Extract MCP_CARD_VISIBLE_TOOL_LIMIT to shared constants * refactor: Support arbitrary authorization header schemes * refactor: Consolidate MCP recommendations dismissal into existing storage key * fix: Use case-insensitive comparison for MCP server ID prefix check * refactor: Centralize MCP server visibility logic and extract recommendations hook * refactor: Cleanup --- .../ChatFormActionAddMcpServersSubmenu.svelte | 2 +- .../ChatFormActionAddSheet.svelte | 10 +- .../ChatMessage/ChatMessage.svelte | 2 +- ...tMessageActionCardPermissionRequest.svelte | 40 ++-- .../app/dialogs/DialogMcpServerAddNew.svelte | 55 +++--- .../DialogMcpServerRecommendations.svelte | 180 +++++++++++++++++ .../src/lib/components/app/dialogs/index.ts | 9 + .../components/app/forms/KeyValuePairs.svelte | 23 ++- .../mcp/McpServerCard/McpServerCard.svelte | 2 +- .../McpServerCard/McpServerCardCompact.svelte | 156 +++++++++++++++ .../McpServerCardEditForm.svelte | 49 +++-- .../components/app/mcp/McpServerForm.svelte | 184 ++++++++++++++---- .../app/mcp/McpServerIdentity.svelte | 14 +- tools/ui/src/lib/components/app/mcp/index.ts | 10 + .../app/settings/SettingsMcpServers.svelte | 2 +- tools/ui/src/lib/constants/index.ts | 1 + tools/ui/src/lib/constants/mcp-form.ts | 2 + .../lib/constants/recommended-mcp-servers.ts | 35 ++++ tools/ui/src/lib/constants/settings-keys.ts | 1 + .../ui/src/lib/constants/settings-registry.ts | 10 +- tools/ui/src/lib/constants/storage.ts | 6 +- .../hooks/use-mcp-recommendations.svelte.ts | 85 ++++++++ .../ui/src/lib/services/migration.service.ts | 60 +++++- .../ui/src/lib/stores/conversations.svelte.ts | 22 +-- tools/ui/src/lib/stores/mcp.svelte.ts | 40 +++- tools/ui/src/lib/types/index.ts | 2 + tools/ui/src/lib/types/mcp.d.ts | 21 +- tools/ui/src/routes/+layout.svelte | 9 + .../components/McpServerFormWrapper.svelte | 37 ++++ .../client/mcp-server-form.svelte.test.ts | 133 +++++++++++++ tools/ui/tests/unit/headers.test.ts | 126 ++++++++++++ .../unit/parse-mcp-server-settings.test.ts | 144 ++++++++++++++ .../unit/recommended-mcp-servers.test.ts | 90 +++++++++ 33 files changed, 1411 insertions(+), 151 deletions(-) create mode 100644 tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte create mode 100644 tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardCompact.svelte create mode 100644 tools/ui/src/lib/constants/recommended-mcp-servers.ts create mode 100644 tools/ui/src/lib/hooks/use-mcp-recommendations.svelte.ts create mode 100644 tools/ui/tests/client/components/McpServerFormWrapper.svelte create mode 100644 tools/ui/tests/client/mcp-server-form.svelte.test.ts create mode 100644 tools/ui/tests/unit/headers.test.ts create mode 100644 tools/ui/tests/unit/parse-mcp-server-settings.test.ts create mode 100644 tools/ui/tests/unit/recommended-mcp-servers.test.ts diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpServersSubmenu.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpServersSubmenu.svelte index dd357d6cd..a75f45f37 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpServersSubmenu.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpServersSubmenu.svelte @@ -18,7 +18,7 @@ let mcpSearchQuery = $state(''); let allMcpServers = $derived(mcpStore.getServersSorted()); - let mcpServers = $derived(allMcpServers.filter((s) => s.enabled)); + let mcpServers = $derived(mcpStore.visibleMcpServers); let hasMcpServers = $derived(mcpServers.length > 0); // let hasAnyMcpServers = $derived(allMcpServers.length > 0); let filteredMcpServers = $derived.by(() => { diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte index 2b708aae5..b67fb267b 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte @@ -74,9 +74,7 @@ const sheetItemRowClass = 'flex w-full items-center justify-between gap-2 rounded-md px-3 py-2 text-left text-sm transition-colors hover:bg-accent'; - function getEnabledMcpServers() { - return mcpStore.getServersSorted().filter((s) => s.enabled); - } + let visibleMcpServers = $derived(mcpStore.visibleMcpServers);
@@ -153,13 +151,13 @@ MCP Servers - {getEnabledMcpServers().length} server{getEnabledMcpServers().length !== 1 ? 's' : ''} + {visibleMcpServers.length} server{visibleMcpServers.length !== 1 ? 's' : ''}
- {#each getEnabledMcpServers() as server (server.id)} + {#each visibleMcpServers as server (server.id)} {@const healthState = mcpStore.getHealthCheckState(server.id)} {@const hasError = healthState.status === HealthCheckStatus.ERROR} {@const displayName = mcpStore.getServerLabel(server)} @@ -202,7 +200,7 @@ {/each} - {#if getEnabledMcpServers().length === 0} + {#if visibleMcpServers.length === 0}
No MCP servers configured
diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte index 7189ce1c7..dadcae0c4 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte @@ -43,7 +43,7 @@ assistantMessages: number; messageTypes: string[]; } | null>(null); - let editedContent = $state(message.content); + let editedContent = $derived(message.content); let rawEditContent = $derived.by(() => { if (message.role !== MessageRole.ASSISTANT) return undefined; diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageActions/ChatMessageActionCard/ChatMessageActionCardPermissionRequest.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageActions/ChatMessageActionCard/ChatMessageActionCardPermissionRequest.svelte index e466c84ee..4337bb6a1 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageActions/ChatMessageActionCard/ChatMessageActionCardPermissionRequest.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageActions/ChatMessageActionCard/ChatMessageActionCardPermissionRequest.svelte @@ -1,8 +1,9 @@ @@ -60,29 +69,27 @@ Add New Server -
- (newServerUrl = v)} - onHeadersChange={(v) => (newServerHeaders = v)} - urlError={newServerUrl ? newServerUrlError : null} - id="new-server" - /> -
+
+
+ (newServerUrl = v)} + onHeadersChange={(v) => (newServerHeaders = v)} + urlError={newServerUrl ? newServerUrlError : null} + id="new-server" + /> +
- - + + - - + + +
diff --git a/tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte b/tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte new file mode 100644 index 000000000..9b4489b82 --- /dev/null +++ b/tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte @@ -0,0 +1,180 @@ + + + + + + Do more with MCP + + Power-up your experience by adding tools, resources and more capabilities provided by MCP + servers. + + + +
+

Quickly get started with

+ + {#each RECOMMENDED_MCP_SERVERS as server (server.id)} + (selected[server.id] = enabled)} + /> + {/each} + + {#if addedServers.length > 0} + {#each addedServers as server (server.id)} + + {/each} + {/if} + + {#if showAddForm} + + (newServerUrl = v)} + onHeadersChange={(v) => (newServerHeaders = v)} + urlError={newServerUrl ? newServerUrlError : null} + id="recommendation-new-server" + /> + +
+ + + +
+
+ {:else} + + + + {/if} +
+ + + + + + +
+
diff --git a/tools/ui/src/lib/components/app/dialogs/index.ts b/tools/ui/src/lib/components/app/dialogs/index.ts index 29136308c..73f22c565 100644 --- a/tools/ui/src/lib/components/app/dialogs/index.ts +++ b/tools/ui/src/lib/components/app/dialogs/index.ts @@ -18,6 +18,15 @@ */ export { default as DialogMcpServerAddNew } from './DialogMcpServerAddNew.svelte'; +/** + * **DialogMcpServerRecommendations** - Suggested MCP servers opt-in dialog + * + * Prompts the user to enable pre-defined recommended MCP servers on first launch. + * Shows one switch per suggested server and persists the choice as a per-chat + * override so the selected servers become available in conversations. + */ +export { default as DialogMcpServerRecommendations } from './DialogMcpServerRecommendations.svelte'; + /** * **DialogExportSettings** - Settings export dialog with sensitive data warning * diff --git a/tools/ui/src/lib/components/app/forms/KeyValuePairs.svelte b/tools/ui/src/lib/components/app/forms/KeyValuePairs.svelte index e0bd8d98e..fd6e59a5b 100644 --- a/tools/ui/src/lib/components/app/forms/KeyValuePairs.svelte +++ b/tools/ui/src/lib/components/app/forms/KeyValuePairs.svelte @@ -1,4 +1,5 @@
@@ -103,6 +123,7 @@ {#each pairs as pair, index (index)}
-
+
{#if showSkeleton} {:else if protocolVersion} diff --git a/tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardCompact.svelte b/tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardCompact.svelte new file mode 100644 index 000000000..6cb3e18b6 --- /dev/null +++ b/tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardCompact.svelte @@ -0,0 +1,156 @@ + + + +
+
+ {#if showSkeleton} + + + + + {:else} + + {/if} +
+ + +
+ + {#if isError && errorMessage} +

{errorMessage}

+ {/if} + + {#if showSkeleton} +
+ +
+ +
+ + + + +
+ {:else} + {#if description} + {#if description.lines === 2} +

+ {description.text} +

+ {:else} +

+ {description.text} +

+ {/if} + {/if} + + {#if tools.length > 0} +
+ {#each visibleTools as tool (tool.name)} + + + + {tool.name} + + + + +

+ {tool.description ?? 'No description'} +

+
+
+ {/each} + + {#if hiddenToolCount > 0} + + + + + {hiddenToolCount} more tools + + + + +

+ {hiddenTools.map((tool) => tool.name).join(', ')} +

+
+
+ {/if} +
+ {/if} + {/if} +
diff --git a/tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardEditForm.svelte b/tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardEditForm.svelte index 6727a9000..8ed4ee8b8 100644 --- a/tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardEditForm.svelte +++ b/tools/ui/src/lib/components/app/mcp/McpServerCard/McpServerCardEditForm.svelte @@ -1,6 +1,7 @@ -
-

Configure Server

+
+
+

Configure Server

- (editUrl = v)} - onHeadersChange={(v) => (editHeaders = v)} - onUseProxyChange={(v) => (editUseProxy = v)} - urlError={editUrl ? urlError : null} - id={serverId} - /> + (editUrl = v)} + onHeadersChange={(v) => (editHeaders = v)} + onUseProxyChange={(v) => (editUseProxy = v)} + urlError={editUrl ? urlError : null} + id={serverId} + /> -
- +
+ - + +
-
+
diff --git a/tools/ui/src/lib/components/app/mcp/McpServerForm.svelte b/tools/ui/src/lib/components/app/mcp/McpServerForm.svelte index 79738e30d..7f05d5fef 100644 --- a/tools/ui/src/lib/components/app/mcp/McpServerForm.svelte +++ b/tools/ui/src/lib/components/app/mcp/McpServerForm.svelte @@ -38,14 +38,87 @@ let headerPairs = $derived(parseHeadersToArray(headers)); + const AUTHORIZATION_HEADER = 'Authorization'; + const BEARER_PREFIX = 'Bearer '; + + // Heuristic: this dedicated UI only owns Authorization headers that already + // carry a Bearer scheme. Anything else (e.g. Basic, raw tokens) stays in the + // KV section so the user can still edit those values verbatim. + const matchesAuthorizationKey = (key: string): boolean => + key.trim().toLowerCase() === AUTHORIZATION_HEADER.toLowerCase(); + + const isBearerScheme = (value: string): boolean => + value.trim().toLowerCase().startsWith(BEARER_PREFIX.toLowerCase()); + + const ownedByBearerUi = (p: KeyValuePair): boolean => + matchesAuthorizationKey(p.key) && isBearerScheme(p.value); + + let hasAuthorization = $derived(headerPairs.some(ownedByBearerUi)); + + let wantsAuthorization = $state(false); + + let showAuthorization = $derived(hasAuthorization || wantsAuthorization); + + let urlInput: HTMLInputElement | null = $state(null); + let bearerInput: HTMLInputElement | null = $state(null); + + $effect(() => { + urlInput?.focus(); + }); + + $effect(() => { + if (wantsAuthorization && bearerInput) { + bearerInput.focus(); + } + }); + + let bearerToken = $derived.by(() => { + const auth = headerPairs.find(ownedByBearerUi); + if (!auth) return ''; + return auth.value.trim().slice(BEARER_PREFIX.length).trim(); + }); + + $effect(() => { + if (!headers.trim()) { + wantsAuthorization = false; + } + }); + function updateHeaderPairs(newPairs: KeyValuePair[]) { headerPairs = newPairs; onHeadersChange(serializeHeaders(newPairs)); } + + // The dedicated UI owns the Authorization slot end-to-end when the user + // engages it: any prior Authorization row (Bearer or otherwise) is replaced + // by exactly one { Authorization: "Bearer " } entry. JSON's last-key + // behavior would otherwise pick one arbitrarily, so we strip first. + function updateBearerToken(token: string) { + const filtered = headerPairs.filter((p) => !matchesAuthorizationKey(p.key)); + + const trimmed = token.trim(); + + if (trimmed) { + filtered.push({ key: AUTHORIZATION_HEADER, value: `${BEARER_PREFIX}${trimmed}` }); + } + + updateHeaderPairs(filtered); + } + + function setUseAuthorization(checked: boolean) { + wantsAuthorization = checked; + + if (!checked) { + // Only drop the entry this UI owns; a non-Bearer Authorization row + // authored in the KV section must survive a toggle off untouched. + const filtered = headerPairs.filter((p) => !ownedByBearerUi(p)); + updateHeaderPairs(filtered); + } + } -
-
+
+
@@ -57,50 +130,52 @@ value={url} oninput={(e) => onUrlChange(e.currentTarget.value)} class={urlError ? 'border-destructive' : ''} + bind:ref={urlInput} /> {#if urlError}

{urlError}

{/if} - - {#if !isWebSocket && onUseProxyChange} - - {/if}
+ + + {#if showAuthorization} +
+ updateBearerToken(e.currentTarget.value)} + class="pl-16" + bind:ref={bearerInput} + /> + + + Bearer + +
+ {/if} + !ownedByBearerUi(p))} + onPairsChange={(pairs) => { + const auth = headerPairs.find(ownedByBearerUi); + updateHeaderPairs(auth ? [...pairs, auth] : pairs); + }} keyPlaceholder="Header name" valuePlaceholder="Value" addButtonLabel="Add" @@ -108,4 +183,37 @@ sectionLabel="Custom Headers" sectionLabelOptional /> + + {#if !isWebSocket && onUseProxyChange} + + {/if}
diff --git a/tools/ui/src/lib/components/app/mcp/McpServerIdentity.svelte b/tools/ui/src/lib/components/app/mcp/McpServerIdentity.svelte index feafc5d81..3f128e02c 100644 --- a/tools/ui/src/lib/components/app/mcp/McpServerIdentity.svelte +++ b/tools/ui/src/lib/components/app/mcp/McpServerIdentity.svelte @@ -1,6 +1,7 @@ + + + {}} + onHeadersChange={(value) => { + headersState = value; + }} + id="mcp-server-form-test" +/> + + diff --git a/tools/ui/tests/client/mcp-server-form.svelte.test.ts b/tools/ui/tests/client/mcp-server-form.svelte.test.ts new file mode 100644 index 000000000..b4bd89234 --- /dev/null +++ b/tools/ui/tests/client/mcp-server-form.svelte.test.ts @@ -0,0 +1,133 @@ +import { describe, expect, it } from 'vitest'; +import { render } from 'vitest-browser-svelte'; +import McpServerFormWrapper from './components/McpServerFormWrapper.svelte'; + +const AUTHORIZATION_HEADER = 'Authorization'; +const BEARER_PREFIX = 'Bearer '; +const BEARER_PLACEHOLDER = 'Paste token here'; + +/** + * Client-side tests for the McpServerForm bearer UI. + * + * The dedicated UI only "owns" Authorization headers that already carry a + * Bearer scheme (heuristic check on the value). Other Authorization values + * stay in the KV section so the user can still edit them verbatim. Storage + * always goes through the same custom-headers slot, so a round-trip via this + * UI produces exactly one `Authorization: Bearer ` entry. + * + * Equivalent parser coverage lives in `tests/unit/headers.test.ts`. + */ +describe('McpServerForm - Authorization / bearer UI', () => { + function bearerInput(screen: Awaited>) { + return screen.locator.getByPlaceholder(BEARER_PLACEHOLDER); + } + + function capturedHeaders(screen: Awaited>) { + return screen.getByTestId('captured-headers'); + } + + it('mounts with the bearer input hidden when no auth header is present', async () => { + const screen = await render(McpServerFormWrapper, { headers: '' }); + + await expect.element(screen.getByRole('textbox', { name: /server url/i })).toBeVisible(); + + await expect.element(bearerInput(screen)).not.toBeInTheDocument(); + }); + + it('toggling Authorization shows the bearer input', async () => { + const screen = await render(McpServerFormWrapper, { headers: '' }); + + await screen.getByRole('switch', { name: /authorization/i }).click(); + + await expect.element(bearerInput(screen)).toBeVisible(); + }); + + it('typing a token writes the Authorization row with the Bearer prefix prepended', async () => { + const screen = await render(McpServerFormWrapper, { headers: '' }); + + await screen.getByRole('switch', { name: /authorization/i }).click(); + + const token = 'super-secret'; + await bearerInput(screen).fill(token); + + const expected = JSON.stringify({ [AUTHORIZATION_HEADER]: `${BEARER_PREFIX}${token}` }); + await expect + .element(capturedHeaders(screen)) + .toHaveAttribute('data-captured-headers', expected); + }); + + it('pre-existing Bearer header pre-fills the bearer input with the token stripped', async () => { + const existing = JSON.stringify({ + 'X-Trace-Id': 'abc', + [AUTHORIZATION_HEADER]: `${BEARER_PREFIX}preexisting` + }); + + const screen = await render(McpServerFormWrapper, { headers: existing }); + + await expect.element(bearerInput(screen)).toBeVisible(); + await expect.element(bearerInput(screen)).toHaveValue('preexisting'); + }); + + it('non-Bearer Authorization is ignored by the dedicated UI and stays in the KV section', async () => { + const existing = JSON.stringify({ [AUTHORIZATION_HEADER]: 'Basic czNjcjpwYXNz' }); + + const screen = await render(McpServerFormWrapper, { headers: existing }); + + await expect.element(bearerInput(screen)).not.toBeInTheDocument(); + + const headerKeyInput = screen.getByPlaceholder('Header name'); + await expect.element(headerKeyInput).toBeVisible(); + }); + + it('engaging the token UI replaces a non-Bearer Authorization with the Bearer scheme', async () => { + const existing = JSON.stringify({ [AUTHORIZATION_HEADER]: 'Basic old' }); + + const screen = await render(McpServerFormWrapper, { headers: existing }); + + await screen.getByRole('switch', { name: /authorization/i }).click(); + await bearerInput(screen).fill('new'); + + const expected = JSON.stringify({ [AUTHORIZATION_HEADER]: `${BEARER_PREFIX}new` }); + await expect + .element(capturedHeaders(screen)) + .toHaveAttribute('data-captured-headers', expected); + }); + + it('toggling Authorization off with no token drops the Bearer row but keeps non-Bearer schemes', async () => { + const existing = JSON.stringify({ [AUTHORIZATION_HEADER]: `${BEARER_PREFIX}xyz` }); + const screen = await render(McpServerFormWrapper, { headers: existing }); + + await screen.getByRole('switch', { name: /authorization/i }).click(); + + await expect.element(capturedHeaders(screen)).toHaveAttribute('data-captured-headers', ''); + }); + + it('toggling Authorization off when no Bearer row is present leaves headers untouched', async () => { + const existing = JSON.stringify({ [AUTHORIZATION_HEADER]: 'Basic czNjcjpwYXNz' }); + const screen = await render(McpServerFormWrapper, { headers: existing }); + + await screen.getByRole('switch', { name: /authorization/i }).click(); + await screen.getByRole('switch', { name: /authorization/i }).click(); + + await expect + .element(capturedHeaders(screen)) + .toHaveAttribute('data-captured-headers', existing); + }); + + it('clearing the bearer input drops the Authorization row', async () => { + const existing = JSON.stringify({ [AUTHORIZATION_HEADER]: `${BEARER_PREFIX}xyz` }); + const screen = await render(McpServerFormWrapper, { headers: existing }); + + await bearerInput(screen).fill(''); + + await expect.element(capturedHeaders(screen)).toHaveAttribute('data-captured-headers', ''); + }); + + it('does not surface Bearer Authorization in the KV section even when pre-existing', async () => { + const existing = JSON.stringify({ [AUTHORIZATION_HEADER]: `${BEARER_PREFIX}xyz` }); + const screen = await render(McpServerFormWrapper, { headers: existing }); + + const headerKeyInput = screen.getByPlaceholder('Header name'); + await expect.element(headerKeyInput).not.toBeInTheDocument(); + }); +}); diff --git a/tools/ui/tests/unit/headers.test.ts b/tools/ui/tests/unit/headers.test.ts new file mode 100644 index 000000000..33619547c --- /dev/null +++ b/tools/ui/tests/unit/headers.test.ts @@ -0,0 +1,126 @@ +import { describe, expect, it } from 'vitest'; +import { parseHeadersToArray, serializeHeaders } from '$lib/utils/headers'; + +/** + * Tests for the header serialization helpers used by the MCP server form + * (custom header rows) and the new Authorization/Bearer-token flow. + */ +describe('parseHeadersToArray', () => { + it('returns an empty array for empty or whitespace-only input', () => { + expect(parseHeadersToArray('')).toEqual([]); + expect(parseHeadersToArray(' ')).toEqual([]); + expect(parseHeadersToArray(undefined as unknown as string)).toEqual([]); + }); + + it('returns an empty array for invalid JSON input', () => { + expect(parseHeadersToArray('{not-json')).toEqual([]); + expect(parseHeadersToArray('[]')).toEqual([]); + expect(parseHeadersToArray('"plain-string"')).toEqual([]); + }); + + it('converts an object into ordered key/value pairs', () => { + expect(parseHeadersToArray('{"X-Foo":"bar","Authorization":"Bearer abc"}')).toEqual([ + { key: 'X-Foo', value: 'bar' }, + { key: 'Authorization', value: 'Bearer abc' } + ]); + }); + + it('stringifies non-string values', () => { + expect(parseHeadersToArray('{"count":"42","flag":"true"}')).toEqual([ + { key: 'count', value: '42' }, + { key: 'flag', value: 'true' } + ]); + }); +}); + +describe('serializeHeaders', () => { + it('returns an empty string when there are no valid pairs', () => { + expect(serializeHeaders([])).toBe(''); + expect(serializeHeaders([{ key: '', value: 'value' }])).toBe(''); + expect(serializeHeaders([{ key: ' ', value: 'value' }])).toBe(''); + }); + + it('returns an empty string when every pair has a blank key', () => { + expect( + serializeHeaders([ + { key: '', value: 'drop-me' }, + { key: ' ', value: 'drop-me-too' }, + { key: '\t', value: 'tab-key' } + ]) + ).toBe(''); + }); + + it('drops pairs with empty keys but keeps the rest', () => { + expect( + serializeHeaders([ + { key: '', value: 'drop-me' }, + { key: 'X-Keep', value: 'ok' } + ]) + ).toBe('{"X-Keep":"ok"}'); + }); + + it('trims keys before serializing', () => { + expect(serializeHeaders([{ key: ' X-Space ', value: 'ok' }])).toBe('{"X-Space":"ok"}'); + }); + + it('preserves the input order of surviving pairs', () => { + const serialized = serializeHeaders([ + { key: 'X-C', value: '3' }, + { key: 'X-A', value: '1' }, + { key: 'X-B', value: '2' } + ]); + + // Object key order follows insertion order in modern JS engines, so + // the serialized JSON writes keys in our input order. + expect(JSON.parse(serialized)).toEqual({ 'X-C': '3', 'X-A': '1', 'X-B': '2' }); + }); +}); + +describe('parseHeadersToArray / serializeHeaders roundtrip', () => { + it('serializes back to an equal header object after a parse', () => { + const original = JSON.stringify({ + 'Content-Type': 'application/json', + 'X-Trace-Id': 'abc-123' + }); + + const roundtrip = serializeHeaders(parseHeadersToArray(original)); + + expect(JSON.parse(roundtrip)).toEqual(JSON.parse(original)); + }); + + it('drops rows whose keys are blank after trimming during serialization', () => { + const pairs = parseHeadersToArray('{"X-Keep":"ok","":"drop-me"}'); + + // parseHeadersToArray keeps raw key strings (the consumer is expected to + // filter blanks, not the parser); serialization must strip them. + expect(pairs).toEqual([ + { key: 'X-Keep', value: 'ok' }, + { key: '', value: 'drop-me' } + ]); + expect(serializeHeaders(pairs)).toBe('{"X-Keep":"ok"}'); + }); + + it('preserves upstream keys untouched (does not lowercase them)', () => { + const upperCased = '{"Authorization":"Bearer xyz"}'; + + const parsed = parseHeadersToArray(upperCased); + + expect(parsed).toEqual([{ key: 'Authorization', value: 'Bearer xyz' }]); + }); + + it('bearer-token write survives a re-parse when paired with regular custom headers', () => { + // The McpServerForm bearer UI writes {Authorization: `Bearer `} + // into the same headers string as the custom KV section. The round + // trip below mirrors the exact shape the form produces so a future + // refactor of either code path cannot silently change the on-disk key. + const pairs = [ + { key: 'X-Trace-Id', value: 'abc-123' }, + { key: 'Authorization', value: 'Bearer super-secret' } + ]; + + const serialized = serializeHeaders(pairs); + + expect(serialized).toBe('{"X-Trace-Id":"abc-123","Authorization":"Bearer super-secret"}'); + expect(parseHeadersToArray(serialized)).toEqual(pairs); + }); +}); diff --git a/tools/ui/tests/unit/parse-mcp-server-settings.test.ts b/tools/ui/tests/unit/parse-mcp-server-settings.test.ts new file mode 100644 index 000000000..956c677d5 --- /dev/null +++ b/tools/ui/tests/unit/parse-mcp-server-settings.test.ts @@ -0,0 +1,144 @@ +import { describe, expect, it, vi } from 'vitest'; +import { parseMcpServerSettings } from '$lib/utils/mcp'; +import { DEFAULT_MCP_CONFIG, MCP_SERVER_ID_PREFIX } from '$lib/constants/mcp'; + +/** + * Tests for the mcpServers settings parser. + * + * The branch seeds the MCP servers setting with a default value of + * `JSON.stringify(RECOMMENDED_MCP_SERVERS)`, so the parser has to be + * resilient to anything that may live in the user's localStorage: malformed + * JSON, wrong shapes, missing fields, falsy-but-not-zero numbers, and entry + * arrays that have been mutated by the user via the settings form. + */ +describe('parseMcpServerSettings', () => { + it('returns an empty array for falsy or whitespace-only input', () => { + expect(parseMcpServerSettings(null)).toEqual([]); + expect(parseMcpServerSettings(undefined)).toEqual([]); + expect(parseMcpServerSettings('')).toEqual([]); + expect(parseMcpServerSettings(' ')).toEqual([]); + }); + + it('returns an empty array and logs a warning for invalid JSON strings', () => { + const warn = vi.spyOn(console, 'warn').mockImplementation(() => {}); + + expect(parseMcpServerSettings('{not-json')).toEqual([]); + expect(warn).toHaveBeenCalled(); + + warn.mockRestore(); + }); + + it('returns an empty array for valid JSON that is not an array', () => { + expect(parseMcpServerSettings('"plain-string"')).toEqual([]); + expect(parseMcpServerSettings('{"id":"foo"}')).toEqual([]); + expect(parseMcpServerSettings('42')).toEqual([]); + expect(parseMcpServerSettings('null')).toEqual([]); + }); + + it('drops entries with no parseable id and substitutes a stable fallback', () => { + const parsed = parseMcpServerSettings( + JSON.stringify([{ url: 'https://a.test', enabled: true }, { url: 'https://b.test' }]) + ); + + expect(parsed).toHaveLength(2); + expect(parsed[0]?.id).toBe(`${MCP_SERVER_ID_PREFIX}-1`); + expect(parsed[1]?.id).toBe(`${MCP_SERVER_ID_PREFIX}-2`); + }); + + it('reuses the first id when it is present and falls back only for missing ones', () => { + const parsed = parseMcpServerSettings( + JSON.stringify([ + { id: 'custom-1', url: 'https://a.test' }, + { url: 'https://b.test' }, + { id: 'custom-3', url: 'https://c.test' } + ]) + ); + + expect(parsed[0]?.id).toBe('custom-1'); + expect(parsed[1]?.id).toBe(`${MCP_SERVER_ID_PREFIX}-2`); + expect(parsed[2]?.id).toBe('custom-3'); + }); + + it('falls back to the configured default requestTimeoutSeconds only for nullish values', () => { + const fallback = DEFAULT_MCP_CONFIG.requestTimeoutSeconds; + + const parsed = parseMcpServerSettings( + JSON.stringify([ + { id: 'a', url: 'https://a.test' }, + { id: 'b', url: 'https://b.test', requestTimeoutSeconds: undefined }, + { id: 'c', url: 'https://c.test', requestTimeoutSeconds: 0 }, + { id: 'd', url: 'https://d.test', requestTimeoutSeconds: 45 } + ]) + ); + + // The parser uses ?? for timeout fallback, which only triggers on + // null/undefined. Explicit 0 is preserved at face value. + expect(parsed[0]?.requestTimeoutSeconds).toBe(fallback); + expect(parsed[1]?.requestTimeoutSeconds).toBe(fallback); + expect(parsed[2]?.requestTimeoutSeconds).toBe(0); + expect(parsed[3]?.requestTimeoutSeconds).toBe(45); + }); + + it('treats whitespace-only headers strings as undefined', () => { + const parsed = parseMcpServerSettings( + JSON.stringify([ + { id: 'a', url: 'https://a.test', headers: ' ' }, + { id: 'b', url: 'https://b.test', headers: '{"X-Foo":"bar"}' } + ]) + ); + + // The parser trims headers and coerces empty/whitespace to undefined. + expect(parsed[0]?.headers).toBeUndefined(); + expect(parsed[1]?.headers).toBe('{"X-Foo":"bar"}'); + }); + + it('defaults coercion for booleans (undefined -> false, true -> true)', () => { + const parsed = parseMcpServerSettings( + JSON.stringify([ + { id: 'a', url: 'https://a.test' }, + { id: 'b', url: 'https://b.test', enabled: true }, + { id: 'c', url: 'https://c.test', enabled: false }, + { id: 'd', url: 'https://d.test', useProxy: true } + ]) + ); + + expect(parsed[0]?.enabled).toBe(false); + expect(parsed[1]?.enabled).toBe(true); + expect(parsed[2]?.enabled).toBe(false); + expect(parsed[0]?.useProxy).toBe(false); + expect(parsed[3]?.useProxy).toBe(true); + }); + + it('preserves input order when mapping entries', () => { + const source = [ + { id: 'gamma', url: 'https://c.test' }, + { id: 'alpha', url: 'https://a.test' }, + { id: 'beta', url: 'https://b.test' } + ]; + + const parsed = parseMcpServerSettings(JSON.stringify(source)); + + expect(parsed.map((entry) => entry.id)).toEqual(['gamma', 'alpha', 'beta']); + }); + + it('passes non-string raw input through the JSON-equality path', () => { + const parsed = parseMcpServerSettings([ + { id: 'a', url: 'https://a.test' }, + { id: 'b', url: 'https://b.test', enabled: true } + ]); + + expect(parsed).toHaveLength(2); + expect(parsed[0]?.id).toBe('a'); + expect(parsed[1]?.enabled).toBe(true); + }); + + it('coerces non-string url values to an empty string rather than throwing', () => { + const parsed = parseMcpServerSettings( + JSON.stringify([{ id: 'a', url: 42 }, { id: 'b' }, { id: 'c', url: 'https://c.test' }]) + ); + + expect(parsed[0]?.url).toBe(''); + expect(parsed[1]?.url).toBe(''); + expect(parsed[2]?.url).toBe('https://c.test'); + }); +}); diff --git a/tools/ui/tests/unit/recommended-mcp-servers.test.ts b/tools/ui/tests/unit/recommended-mcp-servers.test.ts new file mode 100644 index 000000000..3f6fd8f11 --- /dev/null +++ b/tools/ui/tests/unit/recommended-mcp-servers.test.ts @@ -0,0 +1,90 @@ +import { describe, expect, it } from 'vitest'; +import { + RECOMMENDED_MCP_SERVER_IDS, + RECOMMENDED_MCP_SERVERS +} from '$lib/constants/recommended-mcp-servers'; +import { parseMcpServerSettings } from '$lib/utils/mcp'; +import { DEFAULT_MCP_CONFIG, MCP_SERVER_ID_PREFIX } from '$lib/constants/mcp'; + +/** + * Tests for the predefined recommended MCP servers. + * + * These are surfaced to first-time users via + * DialogMcpServerRecommendations and used as the default value of the MCP + * servers setting, so a regression that breaks the round-trip through the + * settings parser would silently break onboarding for new users. + */ +describe('RECOMMENDED_MCP_SERVERS', () => { + it('lists at least one entry and uses stable, unique ids', () => { + expect(RECOMMENDED_MCP_SERVERS.length).toBeGreaterThan(0); + + const ids = RECOMMENDED_MCP_SERVERS.map((server) => server.id); + expect(new Set(ids).size).toBe(ids.length); + + for (const id of ids) { + expect(id).toMatch(/^[a-z0-9-]+$/); + expect(id.toLowerCase()).not.toContain(MCP_SERVER_ID_PREFIX.toLowerCase()); + } + }); + + it('requires a name, description and url for every entry', () => { + for (const server of RECOMMENDED_MCP_SERVERS) { + expect(server.name?.trim().length ?? 0).toBeGreaterThan(0); + expect(server.description.trim().length).toBeGreaterThan(0); + expect(server.url.trim().length).toBeGreaterThan(0); + expect(() => new URL(server.url)).not.toThrow(); + } + }); +}); + +describe('RECOMMENDED_MCP_SERVER_IDS', () => { + it('matches the ids declared in RECOMMENDED_MCP_SERVERS', () => { + expect(RECOMMENDED_MCP_SERVER_IDS.size).toBe(RECOMMENDED_MCP_SERVERS.length); + + for (const server of RECOMMENDED_MCP_SERVERS) { + expect(RECOMMENDED_MCP_SERVER_IDS.has(server.id)).toBe(true); + } + }); +}); + +describe('recommended-mcp-servers default value', () => { + it('round-trips cleanly through parseMcpServerSettings', () => { + const serialized = JSON.stringify(RECOMMENDED_MCP_SERVERS); + const parsed = parseMcpServerSettings(serialized); + + expect(parsed).toHaveLength(RECOMMENDED_MCP_SERVERS.length); + + for (let index = 0; index < RECOMMENDED_MCP_SERVERS.length; index++) { + const source = RECOMMENDED_MCP_SERVERS[index]; + const entry = parsed[index]; + + expect(entry).toBeDefined(); + expect(entry?.id).toBe(source.id); + expect(entry?.url).toBe(source.url); + expect(entry?.enabled).toBe(source.enabled); + expect(entry?.requestTimeoutSeconds).toBe(source.requestTimeoutSeconds); + expect(entry?.name).toBe(source.name); + + // Headers and useProxy are not set on recommended servers; the + // parser must fall back to the inactive defaults rather than + // surfacing undefined-boundary states. + expect(entry?.headers).toBeUndefined(); + expect(entry?.useProxy).toBe(false); + } + }); + + it('uses the global default timeout when one is not specified on an entry', () => { + const sourceOnlyRequired = { + id: 'roundtrip-only', + name: 'Only required fields', + url: 'https://example.test/mcp', + description: 'Smoke entry for parser roundtrip with default timeout.', + enabled: true + }; + + const parsed = parseMcpServerSettings(JSON.stringify([sourceOnlyRequired])); + const entry = parsed[0]; + + expect(entry?.requestTimeoutSeconds).toBe(DEFAULT_MCP_CONFIG.requestTimeoutSeconds); + }); +}); From b5315e16e0c3d19707d4a3e3a9f727f1141d2282 Mon Sep 17 00:00:00 2001 From: Pascal Date: Fri, 3 Jul 2026 12:47:04 +0200 Subject: [PATCH 05/14] server + ui: ping silent SSE streams every 1s and kick only after 3s so slow prefill never drops healthy connections (#25241) * server + ui: ping silent SSE streams every 1s and kick only after 3s so slow prefill never drops healthy connections * server + ui: sse_ping_interval becomes a per-request body field Address review from ngxson: the global default returns to 30 so API clients see no behavior change, and the WebUI sends sse_ping_interval: 1 in the request body since it owns the 3s visibility-kick contract and declares the cadence it needs. Positive values keep the existing > 0 gate, -1 keeps its disabled semantics. * server: move sse_ping_interval into the request schema Address review from ngxson: the field is now a typed field_num with hard limits (-1, INT32_MAX) bound to task_params, seeded from the CLI default alongside the other inherited parameters. The raw json_value read and its redundant comment are gone, and schema evaluation brings type and range validation for free. --- tools/server/README.md | 2 ++ tools/server/server-context.cpp | 9 ++++++--- tools/server/server-schema.cpp | 5 +++++ tools/server/server-task.h | 2 ++ tools/ui/src/lib/constants/stream.ts | 2 +- tools/ui/src/lib/services/chat.service.ts | 1 + tools/ui/src/lib/types/api.d.ts | 1 + 7 files changed, 18 insertions(+), 4 deletions(-) diff --git a/tools/server/README.md b/tools/server/README.md index e88bc5f28..501b66123 100644 --- a/tools/server/README.md +++ b/tools/server/README.md @@ -521,6 +521,8 @@ These words will not be included in the completion, so make sure to add them to `return_progress`: Include prompt processing progress in `stream` mode. The progress will be contained inside `prompt_progress` with 4 values: `total`, `cache`, `processed`, and `time_ms`. The overall progress is `processed/total`, while the actual timed progress is `(processed-cache)/(total-cache)`. The `time_ms` field contains the elapsed time in milliseconds since prompt processing started. Default: `false` +`sse_ping_interval`: Interval in seconds between SSE comment pings emitted while the stream stays silent, keeping the connection observable during long prompt processing. Overrides the server `--sse-ping-interval` setting for this request, `-1` disables pings. Default: server setting + `post_sampling_probs`: Returns the probabilities of top `n_probs` tokens after applying sampling chain. `response_fields`: A list of response fields, for example: `"response_fields": ["content", "generation_settings/n_predict"]`. If the specified field is missing, it will simply be omitted from the response without triggering an error. Note that fields with a slash will be unnested; for example, `generation_settings/n_predict` will move the field `n_predict` from the `generation_settings` object to the root of the response and give it a new name. diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 20e93258f..bb3b91ab5 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -4089,6 +4089,8 @@ std::unique_ptr server_routes::handle_completions_impl( auto & rd = res->rd; auto & params = this->params; + int32_t sse_ping_interval = params.sse_ping_interval; + try { std::vector tasks; @@ -4139,6 +4141,7 @@ std::unique_ptr server_routes::handle_completions_impl( task.params.message_spans = task.tokens.find_message_spans(delimiters); task.id_slot = json_value(data, "id_slot", -1); + sse_ping_interval = task.params.sse_ping_interval; // OAI-compat task.params.res_type = res_type; @@ -4228,7 +4231,7 @@ std::unique_ptr server_routes::handle_completions_impl( } res->status = 200; res->content_type = "text/event-stream"; - res->next = [res_this = res.get(), res_type, &req, ¶ms](std::string & output) -> bool { + res->next = [res_this = res.get(), res_type, sse_ping_interval, &req](std::string & output) -> bool { static auto format_error = [](task_response_type res_type, const json & res_json) { if (res_type == TASK_RESPONSE_TYPE_ANTHROPIC) { return format_anthropic_sse({ @@ -4277,10 +4280,10 @@ std::unique_ptr server_routes::handle_completions_impl( // receive subsequent results bool timeout = false; int64_t start_time = ggml_time_ms(); - auto result = rd.next([&timeout, &start_time, ¶ms, &effective_should_stop]() { + auto result = rd.next([&timeout, &start_time, sse_ping_interval, &effective_should_stop]() { if (effective_should_stop()) { return true; // should_stop condition met - } else if (params.sse_ping_interval > 0 && ggml_time_ms() - start_time > (int64_t)params.sse_ping_interval * 1000) { + } else if (sse_ping_interval > 0 && ggml_time_ms() - start_time > (int64_t)sse_ping_interval * 1000) { timeout = true; return true; // timeout } diff --git a/tools/server/server-schema.cpp b/tools/server/server-schema.cpp index 07a842bd6..5713cc831 100644 --- a/tools/server/server-schema.cpp +++ b/tools/server/server-schema.cpp @@ -37,6 +37,10 @@ std::vector> make_llama_cmpl_schema(const common_params & add((new field_bool("return_progress", params.return_progress)) ->set_desc("Include prompt processing progress events in stream mode")); + add((new field_num("sse_ping_interval", params.sse_ping_interval)) + ->set_hard_limits(-1, INT32_MAX) + ->set_desc("Interval in seconds between SSE comment pings emitted while the stream stays silent, -1 disables pings")); + add((new field_num("n_predict", params.n_predict)) ->set_hard_limits(-1, INT32_MAX) ->add_alias("max_completion_tokens") @@ -504,6 +508,7 @@ task_params eval_llama_cmpl_schema( params.n_cache_reuse = params_base.n_cache_reuse; params.cache_prompt = params_base.cache_prompt; params.antiprompt = params_base.antiprompt; + params.sse_ping_interval = params_base.sse_ping_interval; // enabling this will output extra debug information in the HTTP responses from the server params.verbose = params_base.verbosity > 9; diff --git a/tools/server/server-task.h b/tools/server/server-task.h index 293bdf053..49f62d386 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -54,6 +54,8 @@ struct task_params { bool return_tokens = false; bool return_progress = false; + int32_t sse_ping_interval = 30; // seconds between SSE comment pings while the stream stays silent, -1 disables + int32_t n_keep = 0; // number of tokens to keep from initial prompt int32_t n_discard = 0; // number of tokens after n_keep that may be discarded when shifting context, 0 defaults to half int32_t n_predict = -1; // new tokens to predict diff --git a/tools/ui/src/lib/constants/stream.ts b/tools/ui/src/lib/constants/stream.ts index 3d042451f..67951ee95 100644 --- a/tools/ui/src/lib/constants/stream.ts +++ b/tools/ui/src/lib/constants/stream.ts @@ -1,3 +1,3 @@ // grace window after a visibilitychange before we kick a reader whose socket likely died // while the tab was hidden. covers brief background pauses without thrashing live streams -export const STREAM_VISIBILITY_KICK_MS = 1000; +export const STREAM_VISIBILITY_KICK_MS = 3000; diff --git a/tools/ui/src/lib/services/chat.service.ts b/tools/ui/src/lib/services/chat.service.ts index 7dfee3773..dbd0a94f1 100644 --- a/tools/ui/src/lib/services/chat.service.ts +++ b/tools/ui/src/lib/services/chat.service.ts @@ -255,6 +255,7 @@ export class ChatService { }), stream, return_progress: stream ? true : undefined, + sse_ping_interval: stream ? 1 : undefined, tools: tools && tools.length > 0 ? tools : undefined }; diff --git a/tools/ui/src/lib/types/api.d.ts b/tools/ui/src/lib/types/api.d.ts index ec695ac61..5421f8b7f 100644 --- a/tools/ui/src/lib/types/api.d.ts +++ b/tools/ui/src/lib/types/api.d.ts @@ -265,6 +265,7 @@ export interface ApiChatCompletionRequest { stream?: boolean; model?: string; return_progress?: boolean; + sse_ping_interval?: number; tools?: ApiChatCompletionTool[]; // Reasoning parameters reasoning_format?: string; From 067de937183141f54c681ed684f540706d2c420a Mon Sep 17 00:00:00 2001 From: Pascal Date: Fri, 3 Jul 2026 13:14:52 +0200 Subject: [PATCH 06/14] ui: align persisted config with strict server schema and enable thinking by default (#25242) * ui: migrate legacy string-encoded booleans in persisted config * ui: enable thinking by default Fresh users and legacy conversations without a persisted thinking preference now default to enabled. The per-conversation toggle and the persisted localStorage choice keep taking precedence. Picks up the enable_thinking default from #24876. --- .../ui/src/lib/services/migration.service.ts | 38 ++++++++++++++++++- .../ui/src/lib/stores/conversations.svelte.ts | 11 +++--- 2 files changed, 42 insertions(+), 7 deletions(-) diff --git a/tools/ui/src/lib/services/migration.service.ts b/tools/ui/src/lib/services/migration.service.ts index 152b78cd3..d7709bc6b 100644 --- a/tools/ui/src/lib/services/migration.service.ts +++ b/tools/ui/src/lib/services/migration.service.ts @@ -551,13 +551,49 @@ const mcpDefaultEnabledMigration: Migration = { } }; +const CONFIG_TYPES_MIGRATION_ID = 'config-type-normalization-v1'; + +const configTypesMigration: Migration = { + id: CONFIG_TYPES_MIGRATION_ID, + description: 'Coerce legacy string-encoded booleans in persisted config to real booleans', + + async run(): Promise { + const configRaw = localStorage.getItem(CONFIG_LOCALSTORAGE_KEY); + if (configRaw === null) return; + + const config = JSON.parse(configRaw); + let changed = false; + + // Pre-schema configs persisted booleans as the strings "true"/"false", which the + // strict server schema now rejects. Coerce those back to real booleans. No config + // string field holds exactly "true"/"false", so the match is unambiguous. + for (const key of Object.keys(config)) { + if (config[key] === 'true') { + config[key] = true; + changed = true; + } else if (config[key] === 'false') { + config[key] = false; + changed = true; + } + } + + if (changed) { + localStorage.setItem(CONFIG_LOCALSTORAGE_KEY, JSON.stringify(config)); + } + + if (import.meta.env.DEV && import.meta.env.VITE_DEBUG) + console.log(`[Migration] Config types: coerced string booleans (changed=${changed})`); + } +}; + const migrations: Migration[] = [ localStorageMigration, idxdbMigration, legacyMessageMigration, themeMigration, customJsonKeyMigration, - mcpDefaultEnabledMigration + mcpDefaultEnabledMigration, + configTypesMigration ]; export const MigrationService = { diff --git a/tools/ui/src/lib/stores/conversations.svelte.ts b/tools/ui/src/lib/stores/conversations.svelte.ts index ef8d61309..486202207 100644 --- a/tools/ui/src/lib/stores/conversations.svelte.ts +++ b/tools/ui/src/lib/stores/conversations.svelte.ts @@ -114,14 +114,13 @@ class ConversationsStore { /** Load thinking-enabled default from localStorage */ private static loadThinkingDefaults(): boolean { - if (typeof globalThis.localStorage === 'undefined') return false; + if (typeof globalThis.localStorage === 'undefined') return true; try { const raw = localStorage.getItem(THINKING_ENABLED_DEFAULT_LOCALSTORAGE_KEY); - if (!raw) return false; - const parsed = raw === 'true'; - return typeof parsed === 'boolean' ? parsed : false; + if (!raw) return true; + return raw === 'true'; } catch { - return false; + return true; } } @@ -333,7 +332,7 @@ class ConversationsStore { } this.pendingMcpServerOverrides = []; - this.pendingThinkingEnabled = false; + this.pendingThinkingEnabled = ConversationsStore.loadThinkingDefaults(); this.activeConversation = conversation; if (conversation.currNode) { From 75a48a90559abf65df3f3616a53bb16e5afb9d07 Mon Sep 17 00:00:00 2001 From: "Piotr Wilkin (ilintar)" Date: Fri, 3 Jul 2026 15:36:55 +0200 Subject: [PATCH 07/14] cuda: enable topk-moe fusion for 288 experts (#25267) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * cuda: enable topk-moe fusion for 288 experts The topk-moe fusion only accepted power-of-2 expert counts (or the special-cased 576), so models with 288 experts (e.g. Step-3.7-Flash) fell back to the unfused per-layer routing chain: softmax/sigmoid, argsort, get_rows, sum_rows, div, clamp, scale. At batch size 1 that is ~330 extra tiny graph nodes per token. 288 is a multiple of the warp size, so the existing kernel already handles it; this adds the missing template instantiation and accepts 288 in the eligibility check. Measured on gfx1151 with Step-3.7-Flash IQ4_XS (llama-bench, -b 4096 -ub 4096 -fa 1 -dio 1 -ctk q8_0 -ctv q8_0; machine idle, before/after paired so pp4096 stays matched as a load control): test | before | after ----------------+----------------+---------------- pp4096 | 460.99 ± 0.45 | 462.47 ± 0.34 (unchanged) tg128 | 19.10 ± 0.04 | 19.56 ± 0.03 (+2.4%) tg128 @ d30000 | 12.68 ± 0.04 | 12.69 ± 0.03 (unchanged) Prompt processing is unaffected (the fusion only touches decode routing). The decode gain is ~+2.4% at shallow context and fades with depth: by 30k tokens each step is attention-bound over the KV cache, so removing the fixed routing overhead is no longer visible. Assisted-By: Claude Fable 5 * Update tests/test-backend-ops.cpp Co-authored-by: Oliver Simons * Add comment for case 288 in topk-moe.cu --------- Co-authored-by: Oliver Simons --- ggml/src/ggml-cuda/topk-moe.cu | 8 +++++++- tests/test-backend-ops.cpp | 1 + 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-cuda/topk-moe.cu b/ggml/src/ggml-cuda/topk-moe.cu index c4253bfa4..c80394e31 100644 --- a/ggml/src/ggml-cuda/topk-moe.cu +++ b/ggml/src/ggml-cuda/topk-moe.cu @@ -312,6 +312,10 @@ static void launch_topk_moe_cuda(ggml_backend_cuda_context & ctx, ggml_cuda_kernel_launch(topk_moe_cuda<256, has_bias>, launch_params, logits, weights, ids, bias, n_rows, n_expert_used, clamp_val, scale_val, config); break; + case 288: // StepFun 3.7 + ggml_cuda_kernel_launch(topk_moe_cuda<288, has_bias>, launch_params, + logits, weights, ids, bias, n_rows, n_expert_used, clamp_val, scale_val, config); + break; case 512: ggml_cuda_kernel_launch(topk_moe_cuda<512, has_bias>, launch_params, logits, weights, ids, bias, n_rows, n_expert_used, clamp_val, scale_val, config); @@ -377,8 +381,10 @@ bool ggml_cuda_should_use_topk_moe(const ggml_tensor * gating_op, const ggml_tensor * weights, const ggml_tensor * logits, const ggml_tensor * ids) { + // must match an instantiation of launch_topk_moe_cuda: a power of 2 up to 512, + // or one of the non-power-of-2 expert counts of supported models const int n_expert = ids->nb[1] / ids->nb[0]; - if (((n_expert & (n_expert - 1)) != 0 || n_expert > 512) && n_expert != 576) { + if (((n_expert & (n_expert - 1)) != 0 || n_expert > 512) && n_expert != 288 && n_expert != 576) { return false; } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 34811bd50..5d1f7d8ad 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9219,6 +9219,7 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_topk_moe({128, 1, 1, 1}, 128, with_norm, bias_probs, gate, scale_w)); test_cases.emplace_back(new test_topk_moe({129, 1, 1, 1}, 128, with_norm, bias_probs, gate, scale_w)); test_cases.emplace_back(new test_topk_moe({160, 4, 1, 1}, 160, with_norm, bias_probs, gate, scale_w)); + test_cases.emplace_back(new test_topk_moe({288, 22, 1, 1}, 8, with_norm, bias_probs, gate, scale_w)); // Used by StepFun 3.7 } } } From 152d337fadb93c2a099653c4072d5512c92c5bfd Mon Sep 17 00:00:00 2001 From: Ruixiang Wang Date: Fri, 3 Jul 2026 15:40:06 +0200 Subject: [PATCH 08/14] spec: support spec-draft-p-min in DFlash (#25246) * spec: support spec-draft-p-min in DFlash * dflash: add n_min guard * dflash: guard both n_min and n_max --- common/speculative.cpp | 19 ++++++++++++++----- 1 file changed, 14 insertions(+), 5 deletions(-) diff --git a/common/speculative.cpp b/common/speculative.cpp index 3951bbed5..5b26597f6 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -955,10 +955,11 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl { LOG_INF("%s: - block_size=%d, mask_token_id=%d, n_extract=%u\n", __func__, block_size, mask_token_id, target_layer_ids_n); // DFlash input is [id_last, * (block_size-1)], so it can draft at most block_size-1 tokens per step - if (this->params.n_max > block_size - 1) { - LOG_WRN("%s: requested draft size %d exceeds the trained DFlash block size %d -- clamping to %d draft tokens per step\n", - __func__, this->params.n_max, block_size - 1, block_size - 1); - this->params.n_max = block_size - 1; + if (this->params.n_max > block_size - 1 || this->params.n_min > block_size - 1) { + LOG_WRN("%s: requested draft size (n_max=%d, n_min=%d) exceeds the trained DFlash block size %d -- clamping to %d\n", + __func__, this->params.n_max, this->params.n_min, block_size, block_size - 1); + this->params.n_max = std::min(this->params.n_max, block_size - 1); + this->params.n_min = std::min(this->params.n_min, block_size - 1); } batch = llama_batch_init(llama_n_batch(ctx_dft), 0, n_seq); @@ -968,7 +969,7 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl { for (auto & s : smpls) { common_params_sampling sparams; sparams.no_perf = false; - sparams.top_k = 1; + sparams.top_k = 10; sparams.samplers = { COMMON_SAMPLER_TYPE_TOP_K }; s.reset(common_sampler_init(model_dft, sparams)); } @@ -1173,10 +1174,18 @@ struct common_speculative_impl_draft_dflash : public common_speculative_impl { const llama_token id = cur_p->data[0].id; + if (cur_p->data[0].p < params.p_min) { + break; + } + common_sampler_accept(smpl, id, true); result.push_back(id); } + + if (result.size() < (size_t) params.n_min) { + result.clear(); + } } } From f113e02d5ab4c4910709d46e8d81af92ef945289 Mon Sep 17 00:00:00 2001 From: Pascal Date: Fri, 3 Jul 2026 17:32:48 +0200 Subject: [PATCH 09/14] ui: strip path and weight extension from model id in single model mode (#25137) --- tools/ui/src/lib/constants/model-id.ts | 5 +++++ tools/ui/src/lib/services/models.service.ts | 15 ++++++++++----- 2 files changed, 15 insertions(+), 5 deletions(-) diff --git a/tools/ui/src/lib/constants/model-id.ts b/tools/ui/src/lib/constants/model-id.ts index ee314d167..ab7932240 100644 --- a/tools/ui/src/lib/constants/model-id.ts +++ b/tools/ui/src/lib/constants/model-id.ts @@ -37,3 +37,8 @@ export const MODEL_ACTIVATED_PARAMS_RE = /^[Aa]\d+(\.\d+)?[BbMmKkTt]$/; * Container format segments to exclude from tags (every model uses these). */ export const MODEL_IGNORED_SEGMENTS = new Set(['GGUF', 'GGML']); + +/** + * Matches a trailing weight file extension, e.g. `model.gguf` -> `model`. + */ +export const MODEL_WEIGHT_EXTENSION_RE = /\.(gguf|ggml)$/i; diff --git a/tools/ui/src/lib/services/models.service.ts b/tools/ui/src/lib/services/models.service.ts index 209bd7cab..9574da59e 100644 --- a/tools/ui/src/lib/services/models.service.ts +++ b/tools/ui/src/lib/services/models.service.ts @@ -1,5 +1,5 @@ import { ServerModelStatus } from '$lib/enums'; -import { apiFetch, apiPost } from '$lib/utils'; +import { apiFetch, apiPost, normalizeModelName } from '$lib/utils'; import type { ParsedModelId } from '$lib/types/models'; import { MODEL_QUANTIZATION_SEGMENT_RE, @@ -7,6 +7,7 @@ import { MODEL_PARAMS_RE, MODEL_ACTIVATED_PARAMS_RE, MODEL_IGNORED_SEGMENTS, + MODEL_WEIGHT_EXTENSION_RE, MODEL_ID_NOT_FOUND, MODEL_ID_ORG_SEPARATOR, MODEL_ID_SEGMENT_SEPARATOR, @@ -139,15 +140,19 @@ export class ModelsService { tags: [] }; + // strip directory path and weight extension so a bare `-m /path/file.gguf` + // parses like a clean repo id; the HF `org/model` form is preserved + const source = normalizeModelName(modelId).replace(MODEL_WEIGHT_EXTENSION_RE, ''); + // 1. Extract colon-separated quantization (e.g. `model:Q4_K_M`) - const colonIdx = modelId.indexOf(MODEL_ID_QUANTIZATION_SEPARATOR); + const colonIdx = source.indexOf(MODEL_ID_QUANTIZATION_SEPARATOR); let modelPath: string; if (colonIdx !== MODEL_ID_NOT_FOUND) { - result.quantization = modelId.slice(colonIdx + 1) || null; - modelPath = modelId.slice(0, colonIdx); + result.quantization = source.slice(colonIdx + 1) || null; + modelPath = source.slice(0, colonIdx); } else { - modelPath = modelId; + modelPath = source; } // 2. Extract org name (e.g. `org/model` -> org = "org") From d4cff114c0084f1fbc9b4c62717eca8fb2ae494a Mon Sep 17 00:00:00 2001 From: Nick Towle Date: Fri, 3 Jul 2026 10:03:51 -0700 Subject: [PATCH 10/14] ui: Improve performance when streaming (#25225) * ui: Improve performance when streaming * ui: build sibling info map in branching utils Moves the node map and sibling map construction from the .by block into buildSiblingInfoMap() in branching.ts. The map is built once per structural change and only read afterwards, so it does not need SvelteMap reactivity. Keeping the construction in plain TypeScript fixes the svelte/prefer-svelte-reactivity lint error and groups the branching logic where it already lives. --------- Co-authored-by: Pascal --- .../app/chat/ChatMessages/ChatMessages.svelte | 18 +-- tools/ui/src/lib/utils/branching.ts | 134 +++++------------- tools/ui/src/lib/utils/index.ts | 5 +- 3 files changed, 50 insertions(+), 107 deletions(-) diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessages.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessages.svelte index 2d959cfc2..88efcc4b0 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessages.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessages.svelte @@ -20,9 +20,9 @@ agenticInjectSteeringMessage } from '$lib/stores/agentic.svelte'; import { + buildSiblingInfoMap, copyToClipboard, formatMessageForClipboard, - getMessageSiblings, hasAgenticContent } from '$lib/utils'; @@ -169,6 +169,8 @@ }); }); + let siblingInfoByMessageId = $derived(buildSiblingInfoMap(allConversationMessages)); + let displayMessages = $derived.by(() => { if (!messages.length) { return []; @@ -223,18 +225,18 @@ } } - const siblingInfo = getMessageSiblings(allConversationMessages, msg.id); + const siblingInfo = siblingInfoByMessageId.get(msg.id) ?? { + message: msg, + siblingIds: [msg.id], + currentIndex: 0, + totalSiblings: 1 + }; result.push({ message: msg, toolMessages, isLastAssistantMessage: false, - siblingInfo: siblingInfo || { - message: msg, - siblingIds: [msg.id], - currentIndex: 0, - totalSiblings: 1 - } + siblingInfo }); } diff --git a/tools/ui/src/lib/utils/branching.ts b/tools/ui/src/lib/utils/branching.ts index 4e117b3c2..c40abbdd6 100644 --- a/tools/ui/src/lib/utils/branching.ts +++ b/tools/ui/src/lib/utils/branching.ts @@ -92,18 +92,14 @@ export function filterByLeafNodeId( * Finds the leaf node (message with no children) for a given message branch. * Traverses down the tree following the last child until reaching a leaf. * - * @param messages - All messages in the conversation + * @param nodeMap - Map of messages keyed by ID * @param messageId - Starting message ID to find leaf for * @returns The leaf node ID, or the original messageId if no children */ -export function findLeafNode(messages: readonly DatabaseMessage[], messageId: string): string { - const nodeMap = new Map(); - - // Build node map for quick lookups - for (const msg of messages) { - nodeMap.set(msg.id, msg); - } - +function findLeafNodeInMap( + nodeMap: ReadonlyMap, + messageId: string +): string { let currentNode: DatabaseMessage | undefined = nodeMap.get(messageId); while (currentNode && currentNode.children.length > 0) { // Follow the last child (most recent branch) @@ -114,6 +110,22 @@ export function findLeafNode(messages: readonly DatabaseMessage[], messageId: st return currentNode?.id ?? messageId; } +/** + * Convenience wrapper around {@link findLeafNodeInMap} for callers that only have + * a flat message array. + * + * Finds the leaf node (message with no children) for a given message branch. + * Traverses down the tree following the last child until reaching a leaf. + * + * @param messages - All messages in the conversation + * @param messageId - Starting message ID to find leaf for + * @returns The leaf node ID, or the original messageId if no children + */ +export function findLeafNode(messages: readonly DatabaseMessage[], messageId: string): string { + const nodeMap = new Map(messages.map((msg) => [msg.id, msg] as const)); + return findLeafNodeInMap(nodeMap, messageId); +} + /** * Finds all descendant messages (children, grandchildren, etc.) of a given message. * This is used for cascading deletion to remove all messages in a branch. @@ -156,21 +168,14 @@ export function findDescendantMessages( * Gets sibling information for a message, including all sibling IDs and current position. * Siblings are messages that share the same parent. * - * @param messages - All messages in the conversation + * @param nodeMap - Map of messages keyed by ID * @param messageId - The message to get sibling info for * @returns Sibling information including leaf node IDs for navigation */ export function getMessageSiblings( - messages: readonly DatabaseMessage[], + nodeMap: ReadonlyMap, messageId: string ): ChatMessageSiblingInfo | null { - const nodeMap = new Map(); - - // Build node map for quick lookups - for (const msg of messages) { - nodeMap.set(msg.id, msg); - } - const message = nodeMap.get(messageId); if (!message) { return null; @@ -203,7 +208,9 @@ export function getMessageSiblings( // Convert sibling message IDs to their corresponding leaf node IDs // This allows navigation between different conversation branches - const siblingLeafIds = siblingIds.map((siblingId: string) => findLeafNode(messages, siblingId)); + const siblingLeafIds = siblingIds.map((siblingId: string) => + findLeafNodeInMap(nodeMap, siblingId) + ); // Find current message's position among siblings const currentIndex = siblingIds.indexOf(messageId); @@ -217,85 +224,22 @@ export function getMessageSiblings( } /** - * Creates a display-ready list of messages with sibling information for UI rendering. - * This is the main function used by chat components to render conversation branches. + * Builds sibling information for every message in a conversation. + * A single node map is shared across all lookups for O(1) access. * * @param messages - All messages in the conversation - * @param leafNodeId - Current leaf node being viewed - * @returns Array of messages with sibling navigation info + * @returns Map of message ID to its sibling information */ -export function getMessageDisplayList( - messages: readonly DatabaseMessage[], - leafNodeId: string -): ChatMessageSiblingInfo[] { - // Get the current conversation path - const currentPath = filterByLeafNodeId(messages, leafNodeId, true); - const result: ChatMessageSiblingInfo[] = []; - - // Add sibling info for each message in the current path - for (const message of currentPath) { - if (message.type === 'root') { - continue; // Skip root messages in display - } - - const siblingInfo = getMessageSiblings(messages, message.id); - if (siblingInfo) { - result.push(siblingInfo); +export function buildSiblingInfoMap( + messages: readonly DatabaseMessage[] +): Map { + const nodeMap = new Map(messages.map((msg) => [msg.id, msg] as const)); + const siblingMap = new Map(); + for (const msg of messages) { + const info = getMessageSiblings(nodeMap, msg.id); + if (info) { + siblingMap.set(msg.id, info); } } - - return result; -} - -/** - * Checks if a message has multiple siblings (indicating branching at that point). - * - * @param messages - All messages in the conversation - * @param messageId - The message to check - * @returns True if the message has siblings - */ -export function hasMessageSiblings( - messages: readonly DatabaseMessage[], - messageId: string -): boolean { - const siblingInfo = getMessageSiblings(messages, messageId); - return siblingInfo ? siblingInfo.totalSiblings > 1 : false; -} - -/** - * Gets the next sibling message ID for navigation. - * - * @param messages - All messages in the conversation - * @param messageId - Current message ID - * @returns Next sibling's leaf node ID, or null if at the end - */ -export function getNextSibling( - messages: readonly DatabaseMessage[], - messageId: string -): string | null { - const siblingInfo = getMessageSiblings(messages, messageId); - if (!siblingInfo || siblingInfo.currentIndex >= siblingInfo.totalSiblings - 1) { - return null; - } - - return siblingInfo.siblingIds[siblingInfo.currentIndex + 1]; -} - -/** - * Gets the previous sibling message ID for navigation. - * - * @param messages - All messages in the conversation - * @param messageId - Current message ID - * @returns Previous sibling's leaf node ID, or null if at the beginning - */ -export function getPreviousSibling( - messages: readonly DatabaseMessage[], - messageId: string -): string | null { - const siblingInfo = getMessageSiblings(messages, messageId); - if (!siblingInfo || siblingInfo.currentIndex <= 0) { - return null; - } - - return siblingInfo.siblingIds[siblingInfo.currentIndex - 1]; + return siblingMap; } diff --git a/tools/ui/src/lib/utils/index.ts b/tools/ui/src/lib/utils/index.ts index 61b9932d3..8474691ac 100644 --- a/tools/ui/src/lib/utils/index.ts +++ b/tools/ui/src/lib/utils/index.ts @@ -26,10 +26,7 @@ export { findLeafNode, findDescendantMessages, getMessageSiblings, - getMessageDisplayList, - hasMessageSiblings, - getNextSibling, - getPreviousSibling + buildSiblingInfoMap } from './branching'; // Code From 2d973636e292ee6f75fadcf08d29cb33511f509f Mon Sep 17 00:00:00 2001 From: "Piotr Wilkin (ilintar)" Date: Fri, 3 Jul 2026 23:12:11 +0200 Subject: [PATCH 11/14] chat: trim messages sent to StepFun parser (fixes long reasoning loops) (#25238) * chat: trim messages sent to StepFun parser (fixes long reasoning loops) * add regression test; remove duplicate template * chat: trim StepFun content parts before rendering The StepFun trim workaround ran on the already-rendered messages, where typed content parts have been concatenated into a single string, so the per-part whitespace could no longer be reached. Move the trim ahead of rendering and apply it to content_parts text as well as the string content and reasoning_content. Adds a content-parts regression test. Co-Authored-By: Piotr Wilkin Assisted-By: Claude Fable 5 --------- Co-authored-by: tarruda --- common/chat.cpp | 28 ++++++- .../templates/stepfun-ai-Step-3.5-Flash.jinja | 80 ------------------- tests/test-chat-auto-parser.cpp | 1 - tests/test-chat.cpp | 53 ++++++++++++ 4 files changed, 80 insertions(+), 82 deletions(-) delete mode 100644 models/templates/stepfun-ai-Step-3.5-Flash.jinja diff --git a/common/chat.cpp b/common/chat.cpp index 6da59f4db..22d2ee4a2 100644 --- a/common/chat.cpp +++ b/common/chat.cpp @@ -2378,6 +2378,23 @@ static void func_args_not_string(json & messages) { } } +// Trim leading/trailing whitespace from message contents before rendering. This +// has to run on the messages (not on the rendered JSON) because templates with +// string-only content caps concatenate typed content parts into a single string +// during rendering, after which the per-part whitespace can no longer be reached. +// Both the plain string content and the text of typed content parts are trimmed. +static void trim_all_content(std::vector & messages) { + for (auto & message : messages) { + message.content = trim_whitespace(message.content); + message.reasoning_content = trim_whitespace(message.reasoning_content); + for (auto & part : message.content_parts) { + if (part.type == "text") { + part.text = trim_whitespace(part.text); + } + } + } +} + } // MiniCPM5 format: @@ -2634,7 +2651,16 @@ static common_chat_params common_chat_templates_apply_jinja(const struct common_ params.tools.is_array() && tmpls->template_tool_use ? *tmpls->template_tool_use : *tmpls->template_default; const auto & src = tmpl.source(); const auto & caps = tmpl.original_caps(); - params.messages = render_message_to_json(inputs.messages, tmpl.original_caps()); + std::vector trimmed_messages; + const std::vector * messages_to_render = &inputs.messages; + if (src.find("You have access to the following functions in JSONSchema format") != std::string::npos) { + // StepFun: trim message contents (including typed content parts) before rendering, + // otherwise leftover whitespace drives the model into reasoning loops (issue #24181) + trimmed_messages = inputs.messages; + workaround::trim_all_content(trimmed_messages); + messages_to_render = &trimmed_messages; + } + params.messages = render_message_to_json(*messages_to_render, tmpl.original_caps()); params.tool_choice = inputs.tool_choice; params.reasoning_format = inputs.reasoning_format; params.enable_thinking = inputs.enable_thinking; diff --git a/models/templates/stepfun-ai-Step-3.5-Flash.jinja b/models/templates/stepfun-ai-Step-3.5-Flash.jinja deleted file mode 100644 index c09ea497d..000000000 --- a/models/templates/stepfun-ai-Step-3.5-Flash.jinja +++ /dev/null @@ -1,80 +0,0 @@ -{% macro render_content(content) %}{% if content is none %}{{- '' }}{% elif content is string %}{{- content }}{% elif content is mapping %}{{- content['value'] if 'value' in content else content['text'] }}{% elif content is iterable %}{% for item in content %}{% if item.type == 'text' %}{{- item['value'] if 'value' in item else item['text'] }}{% elif item.type == 'image' %}{% endif %}{% endfor %}{% endif %}{% endmacro %} -{{bos_token}}{%- if tools %} - {{- '<|im_start|>system\n' }} - {%- if messages[0].role == 'system' %} - {{- render_content(messages[0].content) + '\n\n' }} - {%- endif %} - {{- "# Tools\n\nYou have access to the following functions in JSONSchema format:\n\n" }} - {%- for tool in tools %} - {{- "\n" }} - {{- tool | tojson(ensure_ascii=False) }} - {%- endfor %} - {{- "\n\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n\n\n\nvalue_1\n\n\nThis is the value for the second parameter\nthat can span\nmultiple lines\n\n\n\n\n\nReminder:\n- Function calls MUST follow the specified format: an inner \n...\n block must be nested within \n...\n XML tags\n- Required parameters MUST be specified\n<|im_end|>\n" }} -{%- else %} - {%- if messages[0].role == 'system' %} - {{- '<|im_start|>system\n' + render_content(messages[0].content) + '<|im_end|>\n' }} - {%- endif %} -{%- endif %} -{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %} -{%- for message in messages[::-1] %} - {%- set index = (messages|length - 1) - loop.index0 %} - {%- if ns.multi_step_tool and message.role == "user" and render_content(message.content) is string and not(render_content(message.content).startswith('') and render_content(message.content).endswith('')) %} - {%- set ns.multi_step_tool = false %} - {%- set ns.last_query_index = index %} - {%- endif %} -{%- endfor %} -{%- for message in messages %} - {%- set content = render_content(message.content) %} - {%- if (message.role == "user") or (message.role == "system" and not loop.first) %} - {%- set role_name = 'observation' if (message.role == "system" and not loop.first and message.name == 'observation') else message.role %} - {{- '<|im_start|>' + role_name + '\n' + content + '<|im_end|>' + '\n' }} - {%- elif message.role == "assistant" %} - {%- if message.reasoning_content is string %} - {%- set reasoning_content = render_content(message.reasoning_content) %} - {%- else %} - {%- if '' in content %} - {%- set reasoning_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') %} - {%- set content = content.split('')[-1].lstrip('\n') %} - {%- else %} - {%- set reasoning_content = '' %} - {%- endif %} - {%- endif %} - {%- if loop.index0 > ns.last_query_index %} - {{- '<|im_start|>' + message.role + '\n\n' + reasoning_content + '\n\n' + content }} - {%- else %} - {{- '<|im_start|>' + message.role + '\n' + content }} - {%- endif %} - {%- if message.tool_calls %} - {%- for tool_call in message.tool_calls %} - {%- if tool_call.function is defined %} - {%- set tool_call = tool_call.function %} - {%- endif %} - {{- '\n\n' }} - {%- if tool_call.arguments is defined %} - {%- set arguments = tool_call.arguments %} - {%- for args_name, args_value in arguments|items %} - {{- '\n' }} - {%- set args_value = args_value | tojson(ensure_ascii=False) | safe if args_value is mapping or (args_value is sequence and args_value is not string) else args_value | string %} - {{- args_value }} - {{- '\n\n' }} - {%- endfor %} - {%- endif %} - {{- '\n' }} - {%- endfor %} - {%- endif %} - {{- '<|im_end|>\n' }} - {%- elif message.role == "tool" %} - {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %} - {{- '<|im_start|>tool_response\n' }} - {%- endif %} - {{- '' }} - {{- content }} - {{- '' }} - {%- if loop.last or (messages[loop.index0 + 1].role != "tool") %} - {{- '<|im_end|>\n' }} - {%- endif %} - {%- endif %} -{%- endfor %} -{%- if add_generation_prompt %} - {{- '<|im_start|>assistant\n\n' }} -{%- endif %} diff --git a/tests/test-chat-auto-parser.cpp b/tests/test-chat-auto-parser.cpp index 5cc105753..d15fdd2c0 100644 --- a/tests/test-chat-auto-parser.cpp +++ b/tests/test-chat-auto-parser.cpp @@ -1887,7 +1887,6 @@ static void test_role_markers_all_templates(testing & t) { { "Qwen-Qwen3-0.6B.jinja", "<|im_start|>user", "<|im_start|>assistant" }, { "Qwen-QwQ-32B.jinja", "<|im_start|>user", "<|im_start|>assistant" }, { "StepFun3.5-Flash.jinja", "<|im_start|>user", "<|im_start|>assistant" }, - { "stepfun-ai-Step-3.5-Flash.jinja", "<|im_start|>user", "<|im_start|>assistant" }, // DeepSeek family { "deepseek-ai-DeepSeek-R1-Distill-Llama-8B.jinja", "<|User|>", "<|Assistant|>" }, diff --git a/tests/test-chat.cpp b/tests/test-chat.cpp index 5f71e5da6..e1e0a59e6 100644 --- a/tests/test-chat.cpp +++ b/tests/test-chat.cpp @@ -3155,6 +3155,59 @@ static void test_template_output_peg_parsers(bool detailed_debug) { } } } + + { + // StepFun trimming regression test (see https://github.com/ggml-org/llama.cpp/pull/25238) + auto tmpls = read_templates("models/templates/StepFun3.5-Flash.jinja"); + + common_chat_msg message_chatbot = simple_assist_msg("Let me check.\n\n", "I am thinking.\n\n"); + + { + common_chat_templates_inputs inputs; + inputs.messages = { message_chatbot }; + inputs.add_generation_prompt = true; + + auto params = common_chat_templates_apply(tmpls.get(), inputs); + + if (params.prompt.find("Let me check.\n\n") != std::string::npos) { + throw std::runtime_error("StepFun 3.5: content not trimmed"); + } + + if (params.prompt.find("I am thinking.\n\n") != std::string::npos) { + throw std::runtime_error("StepFun 3.5: reasoning_content not trimmed"); + } + } + + { + // Trimming must also reach typed (text) content parts, not just string content + // (see https://github.com/ggml-org/llama.cpp/pull/25238) + common_chat_msg message_parts; + message_parts.role = "user"; + message_parts.content_parts = { + { /* .type = */ "text", /* .text = */ "First part.\n\n" }, + { /* .type = */ "media_marker", /* .text = */ "<__media__>" }, + { /* .type = */ "text", /* .text = */ "Second part.\n\n" }, + }; + + common_chat_templates_inputs inputs; + inputs.messages = { message_parts }; + inputs.add_generation_prompt = true; + + auto params = common_chat_templates_apply(tmpls.get(), inputs); + + if (params.prompt.find("First part.\n\n") != std::string::npos || + params.prompt.find("Second part.\n\n") != std::string::npos) { + throw std::runtime_error("StepFun 3.5: text content parts not trimmed"); + } + + // the trimmed text itself must still be present + if (params.prompt.find("First part.") == std::string::npos || + params.prompt.find("Second part.") == std::string::npos) { + throw std::runtime_error("StepFun 3.5: text content parts missing after trim"); + } + } + } + } { From ef2d770117db45b05aa7ecd1b0acca36370c5470 Mon Sep 17 00:00:00 2001 From: fairydreaming <166155368+fairydreaming@users.noreply.github.com> Date: Sat, 4 Jul 2026 13:37:37 +0200 Subject: [PATCH 12/14] ggml : fix broken CPU concat implementation for quantized types (#25247) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * ggml : fix broken CPU concat implementation for quantized types * tests : concat tests for quantized types --------- Co-authored-by: Stanisław Szymczyk --- ggml/src/ggml-cpu/ops.cpp | 18 +++++++++++++++--- tests/test-backend-ops.cpp | 6 ++++++ 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 6724686b8..c555831ce 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -1913,7 +1913,11 @@ static void ggml_compute_forward_concat_any( GGML_ASSERT(dim >= 0 && dim < 4); int64_t o[4] = {0, 0, 0, 0}; - o[dim] = src0->ne[dim]; + if (dim == 0) { + o[dim] = src0->ne[dim]/ggml_blck_size(src0->type); + } else { + o[dim] = src0->ne[dim]; + } const char * x; @@ -1921,8 +1925,8 @@ static void ggml_compute_forward_concat_any( for (int i3 = 0; i3 < ne3; i3++) { for (int i2 = ith; i2 < ne2; i2 += nth) { for (int i1 = 0; i1 < ne1; i1++) { - for (int i0 = 0; i0 < ne0; i0++) { - if (i0 < ne00 && i1 < ne01 && i2 < ne02 && i3 < ne03) { + for (int i0 = 0; i0 < ne0/ggml_blck_size(dst->type); i0++) { + if (i0 < ne00/ggml_blck_size(src0->type) && i1 < ne01 && i2 < ne02 && i3 < ne03) { x = (const char *)src0->data + (i0 )*nb00 + (i1 )*nb01 + (i2 )*nb02 + (i3 )*nb03; } else { x = (const char *)src1->data + (i0 - o[0])*nb10 + (i1 - o[1])*nb11 + (i2 - o[2])*nb12 + (i3 - o[3])*nb13; @@ -2071,6 +2075,14 @@ void ggml_compute_forward_concat( ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + + if (ggml_is_quantized(src0->type)) { + GGML_ASSERT(ggml_is_contiguous(src0)); + GGML_ASSERT(ggml_is_contiguous(src1)); + GGML_ASSERT(src0->ne[0] % ggml_blck_size(src0->type) == 0); + GGML_ASSERT(src1->ne[0] % ggml_blck_size(src1->type) == 0); + } switch (src0->type) { case GGML_TYPE_F16: diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 5d1f7d8ad..c21675ef5 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -8918,6 +8918,12 @@ static std::vector> make_test_cases_eval() { } } + for (ggml_type type_a : { GGML_TYPE_Q4_0, GGML_TYPE_Q4_1, GGML_TYPE_Q5_0, GGML_TYPE_Q5_1, GGML_TYPE_Q8_0 }) { + for (int dim : { 0, 1, 2, 3, }) { + test_cases.emplace_back(new test_concat(type_a, {128, 12, 13, 14}, dim == 0 ? 256 : 7, dim, 0)); + } + } + for (ggml_sort_order order : {GGML_SORT_ORDER_ASC, GGML_SORT_ORDER_DESC}) { for (uint32_t i = 4; i <= 1024*1024; i *= 2) { test_cases.emplace_back(new test_argsort(GGML_TYPE_F32, {i-1, 1, 1, 1})); From 665892536dfb1b7532161e3182304bd35c33e768 Mon Sep 17 00:00:00 2001 From: Pascal Date: Sat, 4 Jul 2026 16:12:27 +0200 Subject: [PATCH 13/14] ui: add sync blocks so display/behavior settings can be set via --ui-config-file (#25132) * ui: add sync blocks so display/behavior settings can be set via --ui-config-file * ui: remove enable thinking setting --- tools/ui/src/lib/constants/settings-keys.ts | 1 - .../ui/src/lib/constants/settings-registry.ts | 56 +++++++++++++------ 2 files changed, 40 insertions(+), 17 deletions(-) diff --git a/tools/ui/src/lib/constants/settings-keys.ts b/tools/ui/src/lib/constants/settings-keys.ts index 498a79c73..d4d782705 100644 --- a/tools/ui/src/lib/constants/settings-keys.ts +++ b/tools/ui/src/lib/constants/settings-keys.ts @@ -69,7 +69,6 @@ export const SETTINGS_KEYS = { // Developer DISABLE_REASONING_PARSING: 'disableReasoningParsing', EXCLUDE_REASONING_FROM_CONTEXT: 'excludeReasoningFromContext', - ENABLE_THINKING: 'enableThinking', SHOW_RAW_OUTPUT_SWITCH: 'showRawOutputSwitch', // PY_INTERPRETER_ENABLED: 'pyInterpreterEnabled', JS_SANDBOX_ENABLED: 'jsSandboxEnabled', diff --git a/tools/ui/src/lib/constants/settings-registry.ts b/tools/ui/src/lib/constants/settings-registry.ts index 2047d0b43..b314510ab 100644 --- a/tools/ui/src/lib/constants/settings-registry.ts +++ b/tools/ui/src/lib/constants/settings-registry.ts @@ -185,7 +185,11 @@ const SETTINGS_REGISTRY: Record = { defaultValue: false, type: SettingsFieldType.CHECKBOX, section: SETTINGS_SECTION_SLUGS.GENERAL, - isExperimental: true + isExperimental: true, + sync: { + serverKey: SETTINGS_KEYS.TITLE_GENERATION_USE_LLM, + paramType: SyncableParameterType.BOOLEAN + } }, { key: SETTINGS_KEYS.TITLE_GENERATION_PROMPT, @@ -193,7 +197,11 @@ const SETTINGS_REGISTRY: Record = { help: 'Optional template for the title generation prompt. Use {{USER}} for the user message and {{ASSISTANT}} for the assistant message.', defaultValue: TITLE_GENERATION.DEFAULT_PROMPT, type: SettingsFieldType.TEXTAREA, - section: SETTINGS_SECTION_SLUGS.GENERAL + section: SETTINGS_SECTION_SLUGS.GENERAL, + sync: { + serverKey: SETTINGS_KEYS.TITLE_GENERATION_PROMPT, + paramType: SyncableParameterType.STRING + } }, { key: SETTINGS_KEYS.MAX_IMAGE_RESOLUTION, @@ -201,7 +209,11 @@ const SETTINGS_REGISTRY: Record = { help: 'Images larger than this will be resized before sending to server. Set to 0 to disable.', defaultValue: 0, type: SettingsFieldType.INPUT, - section: SETTINGS_SECTION_SLUGS.GENERAL + section: SETTINGS_SECTION_SLUGS.GENERAL, + sync: { + serverKey: SETTINGS_KEYS.MAX_IMAGE_RESOLUTION, + paramType: SyncableParameterType.NUMBER + } } ] }, @@ -385,7 +397,11 @@ const SETTINGS_REGISTRY: Record = { help: 'Display the current build version in the bottom-right corner of the interface.', defaultValue: false, type: SettingsFieldType.CHECKBOX, - section: SETTINGS_SECTION_SLUGS.DISPLAY + section: SETTINGS_SECTION_SLUGS.DISPLAY, + sync: { + serverKey: SETTINGS_KEYS.SHOW_BUILD_VERSION, + paramType: SyncableParameterType.BOOLEAN + } } ] }, @@ -669,7 +685,11 @@ const SETTINGS_REGISTRY: Record = { help: 'After each response, re-submit the conversation to pre-fill the server KV cache. Makes the next turn faster since the prompt is already encoded while you read the response.', defaultValue: false, type: SettingsFieldType.CHECKBOX, - section: SETTINGS_SECTION_SLUGS.DEVELOPER + section: SETTINGS_SECTION_SLUGS.DEVELOPER, + sync: { + serverKey: SETTINGS_KEYS.PRE_ENCODE_CONVERSATION, + paramType: SyncableParameterType.BOOLEAN + } }, { key: SETTINGS_KEYS.DISABLE_REASONING_PARSING, @@ -677,7 +697,11 @@ const SETTINGS_REGISTRY: Record = { help: 'Send reasoning_format=none so the server returns thinking tokens inline instead of extracting them into a separate field.', defaultValue: false, type: SettingsFieldType.CHECKBOX, - section: SETTINGS_SECTION_SLUGS.DEVELOPER + section: SETTINGS_SECTION_SLUGS.DEVELOPER, + sync: { + serverKey: SETTINGS_KEYS.DISABLE_REASONING_PARSING, + paramType: SyncableParameterType.BOOLEAN + } }, { key: SETTINGS_KEYS.EXCLUDE_REASONING_FROM_CONTEXT, @@ -691,14 +715,6 @@ const SETTINGS_REGISTRY: Record = { paramType: SyncableParameterType.BOOLEAN } }, - { - key: SETTINGS_KEYS.ENABLE_THINKING, - label: 'Enable thinking', - help: 'Enable model thinking/reasoning for each request. When off, the model will skip the thinking phase and go straight to the response.', - defaultValue: false, - type: SettingsFieldType.CHECKBOX, - section: SETTINGS_SECTION_SLUGS.DEVELOPER - }, { key: SETTINGS_KEYS.SHOW_RAW_OUTPUT_SWITCH, label: 'Enable raw output toggle', @@ -717,7 +733,11 @@ const SETTINGS_REGISTRY: Record = { help: 'Expose a run_javascript tool to the model. Code runs in a Web Worker inside a sandboxed iframe with an opaque origin, isolated from the WebUI and its API, with a hard timeout.', defaultValue: false, type: SettingsFieldType.CHECKBOX, - section: SETTINGS_SECTION_SLUGS.DEVELOPER + section: SETTINGS_SECTION_SLUGS.DEVELOPER, + sync: { + serverKey: SETTINGS_KEYS.JS_SANDBOX_ENABLED, + paramType: SyncableParameterType.BOOLEAN + } }, { key: SETTINGS_KEYS.CUSTOM_JSON, @@ -753,7 +773,11 @@ const SETTINGS_REGISTRY: Record = { defaultValue: DEFAULT_MCP_CONFIG.requestTimeoutSeconds, type: SettingsFieldType.INPUT, section: SETTINGS_SECTION_SLUGS.MCP, - isPositiveInteger: true + isPositiveInteger: true, + sync: { + serverKey: SETTINGS_KEYS.MCP_REQUEST_TIMEOUT_SECONDS, + paramType: SyncableParameterType.NUMBER + } } ] } From a4107133a634250c8c9d888bc0bc8520dcfd6105 Mon Sep 17 00:00:00 2001 From: liminfei-amd <91481003+liminfei-amd@users.noreply.github.com> Date: Sun, 5 Jul 2026 04:37:38 +0800 Subject: [PATCH 14/14] llama : add guard for K/V rotation input when buffer is unallocated (#25215) llm_graph_input_attn_kv::set_input and llm_graph_input_attn_kv_iswa::set_input call set_input_k_rot / set_input_v_rot whenever the rotation tensor pointer is non-null, but the tensor's buffer can be unallocated (NULL) when a graph only stores K/V without attending -- e.g. DFlash speculative decoding's KV-injection pass. set_input_k_rot then calls ggml_backend_buffer_is_host() on a NULL buffer and aborts with GGML_ASSERT(buffer). Guard the four k_rot/v_rot inputs with the same "&& ->buffer" check that the adjacent kq_mask inputs already use in these two functions. When the buffer is unallocated there is no data to upload, so skipping is correct. Fixes #25191 Signed-off-by: liminfei-amd <91481003+liminfei-amd@users.noreply.github.com> --- src/llama-graph.cpp | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 4c86e43c1..dc41d6690 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -494,11 +494,11 @@ void llm_graph_input_attn_kv::set_input(const llama_ubatch * ubatch) { mctx->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn); } - if (self_k_rot) { + if (self_k_rot && self_k_rot->buffer) { mctx->set_input_k_rot(self_k_rot); } - if (self_v_rot) { + if (self_v_rot && self_v_rot->buffer) { mctx->set_input_v_rot(self_v_rot); } } @@ -592,19 +592,19 @@ void llm_graph_input_attn_kv_iswa::set_input(const llama_ubatch * ubatch) { mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn); } - if (self_k_rot) { + if (self_k_rot && self_k_rot->buffer) { mctx->get_base()->set_input_k_rot(self_k_rot); } - if (self_v_rot) { + if (self_v_rot && self_v_rot->buffer) { mctx->get_base()->set_input_v_rot(self_v_rot); } - if (self_k_rot_swa) { + if (self_k_rot_swa && self_k_rot_swa->buffer) { mctx->get_swa()->set_input_k_rot(self_k_rot_swa); } - if (self_v_rot_swa) { + if (self_v_rot_swa && self_v_rot_swa->buffer) { mctx->get_swa()->set_input_v_rot(self_v_rot_swa); } }