From 520c506958569e51de2bcb9956ae9a8c902457ca Mon Sep 17 00:00:00 2001 From: Jeff lothar Date: Mon, 2 Dec 2024 14:53:04 +0800 Subject: [PATCH 1/6] tidy --- kllm/CMakeLists.txt | 2 + kllm/config/arg.cc | 1 + kllm/core/km_context.cc | 80 ++- kllm/core/km_context.h | 55 +- kllm/core/sampling.cc | 47 -- kllm/core/sampling.h | 3 - kllm/kai/chat_completions.h | 14 +- kllm/kai/completions.h | 15 +- kllm/kai/embeddings.h | 21 +- kllm/kai/infill.cc | 24 + kllm/kai/infill.h | 10 +- kllm/kai/lora.h | 8 +- kllm/kai/rerank.h | 9 +- kllm/kai/slots.cc | 16 + kllm/kai/slots.h | 18 +- kllm/kai/tokenize.h | 20 +- kllm/kai/verified.h | 65 ++ kllm/openai/oai.h | 2 +- kllm/openai/oai_processor.cc | 161 +++-- kllm/openai/oai_processor.h | 27 +- kllm/openai/{request.cc => parser.cc} | 633 +++++++++-------- kllm/openai/{request.h => parser.h} | 76 ++- kllm/proto/builder.cc | 95 +++ kllm/proto/builder.h | 944 ++++++++++++++++++++++++++ kllm/proto/interface.struct.proto | 69 +- kllm/tools/embedding/emb.cc | 19 +- kllm/utility/fs.cc | 2 +- 27 files changed, 1840 insertions(+), 596 deletions(-) create mode 100644 kllm/kai/verified.h rename kllm/openai/{request.cc => parser.cc} (38%) rename kllm/openai/{request.h => parser.h} (30%) create mode 100644 kllm/proto/builder.cc create mode 100644 kllm/proto/builder.h diff --git a/kllm/CMakeLists.txt b/kllm/CMakeLists.txt index 60eddde..a0c5a4b 100644 --- a/kllm/CMakeLists.txt +++ b/kllm/CMakeLists.txt @@ -62,6 +62,7 @@ file(GLOB_RECURSE CORE_SRC "core/*.cc") file(GLOB_RECURSE OPAI_SRC "openai/*.cc") file(GLOB_RECURSE KAI_SRC "kai/*.cc") file(GLOB_RECURSE CONFIG_SRC "config/*.cc") +file(GLOB_RECURSE PROTO_SRC "proto/*.cc") kmcmake_cc_library( NAMESPACE ${PROJECT_NAME} NAME core @@ -73,6 +74,7 @@ kmcmake_cc_library( ${KAI_SRC} ${TARGET_SRCS} ${CONFIG_SRC} + ${PROTO_SRC} OBJECTS proto_obj CXXOPTS diff --git a/kllm/config/arg.cc b/kllm/config/arg.cc index 7540b6c..fbdea90 100644 --- a/kllm/config/arg.cc +++ b/kllm/config/arg.cc @@ -33,6 +33,7 @@ #include #include #include +#include namespace kllm { diff --git a/kllm/core/km_context.cc b/kllm/core/km_context.cc index ac76412..a985849 100644 --- a/kllm/core/km_context.cc +++ b/kllm/core/km_context.cc @@ -768,7 +768,6 @@ namespace kllm { TaskResult result = queue_results.recv(id_tasks); if (!result_handler(result)) { cancel_tasks(id_tasks); - LOG(WARNING)<<1; break; } @@ -1679,21 +1678,22 @@ namespace kllm { return res >= 0; } - std::string KMContext::chat_apply_template(const std::string & tmpl,const std::vector & msgs,bool add_ass) const { - int alloc_size = 0; - bool fallback = false; // indicate if we must fallback to default chatml - std::vector chat; - for (auto & msg : msgs) { - chat.push_back({msg.role.c_str(), msg.content.c_str()}); - alloc_size += (msg.role.size() + msg.content.size()) * 1.25; - } + std::string KMContext::chat_apply_template(const std::string & tmpl,const std::vector & msgs,bool add_ass) const { + return chat_apply_template(tmpl, msgs.begin(), msgs.end(), add_ass); + } + + std::string KMContext::chat_apply_template(const std::string & tmpl, const ::google::protobuf::RepeatedPtrField & msgs, bool add_ass) const { + return chat_apply_template(tmpl, msgs.begin(), msgs.end(), add_ass); + } + std::string KMContext::chat_apply_template_impl(const std::string & tmpl, const std::vector &chat,bool add_ass, size_t alloc_size) const { + bool fallback = false; // indicate if we must fallback to default chatml const char * ptr_tmpl = tmpl.empty() ? nullptr : tmpl.c_str(); + std::vector buf(alloc_size); // run the first time to get the total output length int32_t res = llama_chat_apply_template(km_model.model, ptr_tmpl, chat.data(), chat.size(), add_ass, buf.data(), buf.size()); - // error: chat template is not supported if (res < 0) { if (ptr_tmpl != nullptr) { @@ -1720,13 +1720,17 @@ namespace kllm { return formatted_chat; } + const std::string &KMContext::chat_template() const { + return params.chat_template; + } + std::string KMContext::chat_format_single(const std::string & tmpl, - const std::vector & past_msg, - const common_chat_msg & new_msg, + const std::vector & past_msg, + const ChatMessage & new_msg, bool add_ass) const { std::ostringstream ss; auto fmt_past_msg = past_msg.empty() ? "" : chat_apply_template(tmpl, past_msg, false); - std::vector chat_new(past_msg); + std::vector chat_new(past_msg); // if the past_msg ends with a newline, we must preserve it in the formatted version if (add_ass && !fmt_past_msg.empty() && fmt_past_msg.back() == '\n') { ss << "\n"; @@ -1739,14 +1743,50 @@ namespace kllm { return ss.str(); } - std::string KMContext::chat_format_example(const std::string & tmpl) const { - std::vector msgs = { - {"system", "You are a helpful assistant"}, - {"user", "Hello"}, - {"assistant", "Hi there"}, - {"user", "How are you?"}, + std::string KMContext::chat_format_single(const std::string & tmpl, + const ::google::protobuf::RepeatedPtrField &past_msg, + const ChatMessage & new_msg, + bool add_ass) const { + std::ostringstream ss; + auto fmt_past_msg = past_msg.empty() ? "" : chat_apply_template(tmpl, past_msg, false); + std::vector chat_new; + for(auto &it : past_msg) { + chat_new.push_back(it); + } + // if the past_msg ends with a newline, we must preserve it in the formatted version + if (add_ass && !fmt_past_msg.empty() && fmt_past_msg.back() == '\n') { + ss << "\n"; }; - return chat_apply_template(tmpl, msgs, true); + // format chat with new_msg + chat_new.push_back(new_msg); + auto fmt_new_msg = chat_apply_template(tmpl, chat_new, add_ass); + // get the diff part + ss << fmt_new_msg.substr(fmt_past_msg.size(), fmt_new_msg.size() - fmt_past_msg.size()); + return ss.str(); + } + + std::once_flag init_chat; + static std::vector msgs; + + void init_once() { + msgs.resize(4); + msgs[0].set_role("system"); + msgs[0].set_content("You are a helpful assistant"); + msgs[1].set_role("user"); + msgs[1].set_content("Hello"); + msgs[2].set_role("assistant"); + msgs[2].set_content("Hi there"); + msgs[3].set_role("user"); + msgs[3].set_content("How are you?"); + } + + const std::vector& get_default_ex() { + std::call_once(init_chat, init_once); + return msgs; + } + std::string KMContext::chat_format_example(const std::string & tmpl) const { + const auto &mgs = get_default_ex(); + return chat_apply_template(tmpl, mgs, true); } void KMContext::string_process_escapes(std::string & input) { diff --git a/kllm/core/km_context.h b/kllm/core/km_context.h index 0e9d984..312b7b0 100644 --- a/kllm/core/km_context.h +++ b/kllm/core/km_context.h @@ -81,7 +81,8 @@ namespace kllm { // // break the input "prompt" into multiple tasks if needed, then format and tokenize the input prompt(s) - std::vector create_tasks_inference(const KaiRequest &data, ServerTaskInfType inf_type, const KaiPrompts&prompts); + std::vector + create_tasks_inference(const KaiRequest &data, ServerTaskInfType inf_type, const KaiPrompts &prompts); std::vector> tokenize_input_prompts(const KaiPrompts &prompt, bool add_special, bool parse_special, @@ -103,6 +104,7 @@ namespace kllm { const std::unordered_set &id_tasks, const std::function &result_handler, const std::function &error_handler); + // // Functions to process the task // @@ -129,24 +131,36 @@ namespace kllm { // Check if the template supplied via "--chat-template" is supported or not. Returns true if it's valid - static bool chat_verify_template(const std::string & tmpl); + static bool chat_verify_template(const std::string &tmpl); + + static void string_process_escapes(std::string &input); - static void string_process_escapes(std::string & input); + std::string string_from(const struct llama_batch &batch) const; - std::string string_from(const struct llama_batch & batch) const; + std::string string_from(const std::vector &tokens) const; - std::string string_from(const std::vector & tokens) const; + const std::string &chat_template() const; // CPP wrapper for llama_chat_apply_template // If the built-in template is not supported, we default to chatml // If the custom "tmpl" is not supported, we throw an error - std::string chat_apply_template(const std::string & tmpl, const std::vector & msgs, bool add_ass) const; + std::string + chat_apply_template(const std::string &tmpl, const std::vector &msgs, bool add_ass) const; + + std::string + chat_apply_template(const std::string &tmpl, const ::google::protobuf::RepeatedPtrField &msgs, + bool add_ass) const; // Format single message, while taking into account the position of that message in chat history - std::string chat_format_single(const std::string & tmpl, - const std::vector & past_msg, - const common_chat_msg & new_msg, - bool add_ass) const; + std::string chat_format_single(const std::string &tmpl, + const std::vector &past_msg, + const ChatMessage &new_msg, + bool add_ass) const; + + std::string chat_format_single(const std::string &tmpl, + const ::google::protobuf::RepeatedPtrField &past_msg, + const ChatMessage &new_msg, + bool add_ass) const; // Returns an example of formatted chat std::string chat_format_example(const std::string &tmpl) const; @@ -159,6 +173,27 @@ namespace kllm { KaiSlotState trans_proto(const server_slot &slot) const; private: + + template + std::string chat_apply_template(const std::string &tmpl, const Itr &begin, const Itr &end, bool add_ass) const { + size_t alloc_size = 0; + std::vector chat; + auto it = begin; + int i = 0; + while (it != end) { + chat.push_back({it->role().c_str(), it->content().c_str()}); + alloc_size += (it->role().size() + it->content().size()) * 1.25; + ++it; + ++i; + } + auto s = chat_apply_template_impl(tmpl, chat, add_ass, alloc_size); + return s; + } + + std::string + chat_apply_template_impl(const std::string &tmpl, const std::vector &chat, bool add_ass, + size_t alloc_size) const; + void send_partial_response(server_slot &slot, completion_token_output tkn); void send_final_response(const server_slot &slot); diff --git a/kllm/core/sampling.cc b/kllm/core/sampling.cc index b672764..571ec5c 100644 --- a/kllm/core/sampling.cc +++ b/kllm/core/sampling.cc @@ -428,53 +428,6 @@ namespace kllm { } } - std::vector - common_sampler_types_from_names(const std::vector &names, bool allow_alt_names) { - std::unordered_map sampler_canonical_name_map{ - {"dry", COMMON_SAMPLER_TYPE_DRY}, - {"top_k", COMMON_SAMPLER_TYPE_TOP_K}, - {"top_p", COMMON_SAMPLER_TYPE_TOP_P}, - {"typ_p", COMMON_SAMPLER_TYPE_TYPICAL_P}, - {"min_p", COMMON_SAMPLER_TYPE_MIN_P}, - {"temperature", COMMON_SAMPLER_TYPE_TEMPERATURE}, - {"xtc", COMMON_SAMPLER_TYPE_XTC}, - {"infill", COMMON_SAMPLER_TYPE_INFILL}, - }; - - // since samplers names are written multiple ways - // make it ready for both system names and input names - std::unordered_map sampler_alt_name_map{ - {"top-k", COMMON_SAMPLER_TYPE_TOP_K}, - {"top-p", COMMON_SAMPLER_TYPE_TOP_P}, - {"nucleus", COMMON_SAMPLER_TYPE_TOP_P}, - {"typical-p", COMMON_SAMPLER_TYPE_TYPICAL_P}, - {"typical", COMMON_SAMPLER_TYPE_TYPICAL_P}, - {"typ-p", COMMON_SAMPLER_TYPE_TYPICAL_P}, - {"typ", COMMON_SAMPLER_TYPE_TYPICAL_P}, - {"min-p", COMMON_SAMPLER_TYPE_MIN_P}, - {"temp", COMMON_SAMPLER_TYPE_TEMPERATURE}, - }; - - std::vector samplers; - samplers.reserve(names.size()); - - for (const auto &name: names) { - auto sampler = sampler_canonical_name_map.find(name); - if (sampler != sampler_canonical_name_map.end()) { - samplers.push_back(sampler->second); - } else { - if (allow_alt_names) { - sampler = sampler_alt_name_map.find(name); - if (sampler != sampler_alt_name_map.end()) { - samplers.push_back(sampler->second); - } - } - } - } - - return samplers; - } - std::vector common_sampler_types_from_chars(const std::string &chars) { std::unordered_map sampler_name_map = { {common_sampler_type_to_chr(COMMON_SAMPLER_TYPE_DRY), COMMON_SAMPLER_TYPE_DRY}, diff --git a/kllm/core/sampling.h b/kllm/core/sampling.h index 3db44a5..5a0613b 100644 --- a/kllm/core/sampling.h +++ b/kllm/core/sampling.h @@ -104,8 +104,5 @@ namespace kllm { std::string common_sampler_type_to_str(KaiSamplerType cnstr); - std::vector - common_sampler_types_from_names(const std::vector &names, bool allow_alt_names); - std::vector common_sampler_types_from_chars(const std::string &chars); } // namespace kllm diff --git a/kllm/kai/chat_completions.h b/kllm/kai/chat_completions.h index d070c8d..8d1d8bd 100644 --- a/kllm/kai/chat_completions.h +++ b/kllm/kai/chat_completions.h @@ -20,12 +20,15 @@ #include #include +#include +#include +#include namespace kllm { - class KaiChatCompletions { + class KaiChatCompletions : public VerifiedRequest{ public: - explicit KaiChatCompletions(turbo::Nonnull c) : _context(c) {} + explicit KaiChatCompletions(turbo::Nonnull c) : VerifiedRequest(c) {} ~KaiChatCompletions() = default; @@ -54,7 +57,10 @@ namespace kllm { return "chatcmpl-" + random_string(); } - private: - KMContext *_context{nullptr}; + turbo::Status verify_impl(KaiRequest &req) override { + const auto formatted_chat = _context->chat_apply_template(_context-> params.chat_template, req.chat_completion_request().segments(), true); + req.mutable_chat_completion_request()->set_message(formatted_chat); + return turbo::OkStatus(); + } }; } // namespace kllm diff --git a/kllm/kai/completions.h b/kllm/kai/completions.h index 9d06174..9f5948b 100644 --- a/kllm/kai/completions.h +++ b/kllm/kai/completions.h @@ -20,12 +20,13 @@ #include #include +#include namespace kllm { - class KaiCompletions { + class KaiCompletions : public VerifiedRequest { public: - explicit KaiCompletions(turbo::Nonnull c) : _context(c) {} + explicit KaiCompletions(turbo::Nonnull c) : VerifiedRequest(c) {} ~KaiCompletions() = default; @@ -35,7 +36,13 @@ namespace kllm { const std::function &func, const std::function &on_complete); - private: - KMContext *_context{nullptr}; + turbo::Status verify_impl(KaiRequest &req) override { + if (req.completion_request().prompts().empty()) { + return turbo::invalid_argument_error("\"prompt\" must be provided"); + } + return turbo::OkStatus(); + } + }; + } // namespace kllm diff --git a/kllm/kai/embeddings.h b/kllm/kai/embeddings.h index ea3a43b..5fe3075 100644 --- a/kllm/kai/embeddings.h +++ b/kllm/kai/embeddings.h @@ -20,19 +20,26 @@ #include #include - +#include namespace kllm { - class KaiEmbeddings { + class KaiEmbeddings : public VerifiedRequest { public: - KaiEmbeddings(turbo::Nonnull c) : _context(c) {} + explicit KaiEmbeddings(turbo::Nonnull c) : VerifiedRequest(c) {} - ~KaiEmbeddings() = default; + ~KaiEmbeddings() override = default; - void embedding(const KaiRequest &req, KaiResponse&res); + void embedding(const KaiRequest &req, KaiResponse &res); - private: - KMContext *_context{nullptr}; + protected: + turbo::Status verify_impl(KaiRequest &req) override { + if (!req.embedding_request().has_prompts() || req.embedding_request().prompts().values().empty()) { + return turbo::invalid_argument_error("\"input\" or \"content\"must be provided"); + } + return turbo::OkStatus(); + } }; + + } // namespace kllm diff --git a/kllm/kai/infill.cc b/kllm/kai/infill.cc index f27a003..4b00e45 100644 --- a/kllm/kai/infill.cc +++ b/kllm/kai/infill.cc @@ -77,4 +77,28 @@ namespace kllm { _context->queue_results.remove_waiting_task_ids(task_ids); on_complete(); } + + turbo::Status KaiInfill::verify_impl(KaiRequest &req) { + if (_context->params.embedding) { + return turbo::unavailable_error( "This server does not support completions. Start it without `--embeddings`"); + } + // check model compatibility + std::string err; + if (llama_token_fim_pre(_context->km_model.model) == LLAMA_TOKEN_NULL) { + err += "prefix token is missing. "; + } + if (llama_token_fim_suf(_context->km_model.model) == LLAMA_TOKEN_NULL) { + err += "suffix token is missing. "; + } + if (llama_token_fim_mid(_context->km_model.model) == LLAMA_TOKEN_NULL) { + err += "middle token is missing. "; + } + if (!err.empty()) { + return turbo::unavailable_error(turbo::str_format("Infill is not supported by this model: %s", err.c_str())); + } + if(!req.infill_request().has_prompts() || req.infill_request().prompts().empty()) { + return turbo::invalid_argument_error("no prompts"); + } + return turbo::OkStatus(); + } } // namespace kllm diff --git a/kllm/kai/infill.h b/kllm/kai/infill.h index 95d0418..6edb0d0 100644 --- a/kllm/kai/infill.h +++ b/kllm/kai/infill.h @@ -20,12 +20,14 @@ #include #include +#include +#include namespace kllm { - class KaiInfill { + class KaiInfill: public VerifiedRequest { public: - explicit KaiInfill(turbo::Nonnull c) : _context(c) {} + explicit KaiInfill(turbo::Nonnull c) : VerifiedRequest(c) {} ~KaiInfill() = default; @@ -35,7 +37,7 @@ namespace kllm { const std::function &func, const std::function &on_complete); - private: - KMContext *_context{nullptr}; + turbo::Status verify_impl(KaiRequest &req) override; + }; } // namespace kllm diff --git a/kllm/kai/lora.h b/kllm/kai/lora.h index 6ecb555..e91cdb5 100644 --- a/kllm/kai/lora.h +++ b/kllm/kai/lora.h @@ -20,22 +20,20 @@ #include #include +#include namespace kllm { - class KaiLora { + class KaiLora: public VerifiedRequest { public: - explicit KaiLora(turbo::Nonnull c) : _context(c) {} + explicit KaiLora(turbo::Nonnull c) : VerifiedRequest(c) {} ~KaiLora() = default; void lora(const KaiRequest &req, KaiResponse&res) const; void lora_apply(const KaiRequest &req, KaiResponse&res); - - private: - KMContext *_context{nullptr}; }; } // namespace kllm diff --git a/kllm/kai/rerank.h b/kllm/kai/rerank.h index cdd9a27..05d78a3 100644 --- a/kllm/kai/rerank.h +++ b/kllm/kai/rerank.h @@ -20,18 +20,17 @@ #include #include +#include +#include namespace kllm { - class KaiRerank { + class KaiRerank: public VerifiedRequest { public: - explicit KaiRerank(turbo::Nonnull c) : _context(c) {} + explicit KaiRerank(turbo::Nonnull c) : VerifiedRequest(c) {} ~KaiRerank() = default; void rank(const KaiRequest &req, KaiResponse &res); - - private: - KMContext *_context{nullptr}; }; } // namespace kllm diff --git a/kllm/kai/slots.cc b/kllm/kai/slots.cc index 0d9aa88..9727e96 100644 --- a/kllm/kai/slots.cc +++ b/kllm/kai/slots.cc @@ -20,6 +20,22 @@ namespace kllm { + void KaiSlots::action(const KaiRequest &req, KaiResponse &res) { + if(req.query_type() == QUERY_SLOTS_SAVE) { + save(req, res); + return; + } else if(req.query_type() == QUERY_SLOTS_RESTORE) { + restore(req, res); + return; + } else if(req.query_type() == QUERY_SLOTS_ERASE) { + erase(req, res); + return; + } + + res.mutable_status()->set_code(ERROR_TYPE_INVALID_REQUEST); + res.mutable_status()->set_errmsg("Invalid action"); + } + void KaiSlots::save(const KaiRequest &req, KaiResponse&res) { std::string filename = req.slots_task().filename(); if (!fs_validate_filename(filename)) { diff --git a/kllm/kai/slots.h b/kllm/kai/slots.h index 34f4c6d..eb38e8f 100644 --- a/kllm/kai/slots.h +++ b/kllm/kai/slots.h @@ -20,14 +20,18 @@ #include #include +#include +#include namespace kllm { - class KaiSlots { + class KaiSlots: public VerifiedRequest { public: - explicit KaiSlots(turbo::Nonnull c) : _context(c) {} + explicit KaiSlots(turbo::Nonnull c) : VerifiedRequest(c) {} - ~KaiSlots() = default; + ~KaiSlots() override = default; + + void action(const KaiRequest &req, KaiResponse &res); void save(const KaiRequest &req, KaiResponse &res); @@ -35,7 +39,11 @@ namespace kllm { void erase(const KaiRequest &req, KaiResponse &res); - private: - KMContext *_context{nullptr}; + turbo::Status verify_impl(KaiRequest &req) override { + std::string filepath = _context->params.slot_save_path + req.slots_task().filename(); + req.mutable_slots_task()->set_filepath(filepath); + return turbo::OkStatus(); + } }; + } // namespace kllm diff --git a/kllm/kai/tokenize.h b/kllm/kai/tokenize.h index 4392d95..dc42293 100644 --- a/kllm/kai/tokenize.h +++ b/kllm/kai/tokenize.h @@ -20,21 +20,27 @@ #include #include +#include namespace kllm { - class KaiTokenize { + class KaiTokenize : public VerifiedRequest { public: - explicit KaiTokenize(turbo::Nonnull c) : _context(c) {} + explicit KaiTokenize(turbo::Nonnull c) : VerifiedRequest(c) {} - ~KaiTokenize() = default; + ~KaiTokenize() override = default; - void tokenize(const KaiRequest &req, KaiResponse&res); - void detokenize(const KaiRequest &req, KaiResponse&res); + void tokenize(const KaiRequest &req, KaiResponse &res); + void detokenize(const KaiRequest &req, KaiResponse &res); + + turbo::Status verify_impl(KaiRequest &req) override { + if(!req.tokenize_request().has_prompts() || req.tokenize_request().prompts().values().empty()) { + return turbo::invalid_argument_error("\"input\" or \"content\"must be provided"); + } + return turbo::OkStatus(); + } - private: - KMContext *_context{nullptr}; }; } \ No newline at end of file diff --git a/kllm/kai/verified.h b/kllm/kai/verified.h new file mode 100644 index 0000000..447e40c --- /dev/null +++ b/kllm/kai/verified.h @@ -0,0 +1,65 @@ +// Copyright (C) 2024 Kumo inc. +// Author: Jeff.li lijippy@163.com +// All rights reserved. +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published +// by the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// + +#pragma once + +#include +#include +#include + +namespace kllm { + + class VerifiedRequest { + public: + explicit VerifiedRequest(turbo::Nonnull c) : _context(c){} + virtual ~VerifiedRequest() = default; + + turbo::Status verify(KaiRequest &req) { + static InternalSamplerParams isp; + auto rs = verify_impl(req); + if(!rs.ok()) { + return rs; + } + if(req.sparam().dry_sequence_breakers().empty()) { + req.mutable_sparam()->mutable_dry_sequence_breakers()->Assign(isp.dry_sequence_breakers.begin(), + isp.dry_sequence_breakers.end()); + req.mutable_sparam()->mutable_samplers()->Assign(isp.samplers.begin(), isp.samplers.end()); + } + + auto &default_sparams = _context->params.sparams; + if(req.sparam().samplers().empty()) { + req.mutable_sparam()->mutable_samplers()->Assign(default_sparams.samplers.begin(), + default_sparams.samplers.end()); + } + if(req.sparam().grammar().empty()) { + req.mutable_sparam()->set_grammar(default_sparams.grammar); + } + + if(req.model().empty()) { + req.set_model(_context->params.model_alias); + } + return turbo::OkStatus(); + } + protected: + virtual turbo::Status verify_impl(KaiRequest &req) { + return turbo::OkStatus(); + } + protected: + KMContext *_context{nullptr}; + }; +} // namespace kllm diff --git a/kllm/openai/oai.h b/kllm/openai/oai.h index dfbe9a7..369e12f 100644 --- a/kllm/openai/oai.h +++ b/kllm/openai/oai.h @@ -18,7 +18,7 @@ #include #include #include -#include +#include #include #include #include diff --git a/kllm/openai/oai_processor.cc b/kllm/openai/oai_processor.cc index 47e84d1..6180dd2 100644 --- a/kllm/openai/oai_processor.cc +++ b/kllm/openai/oai_processor.cc @@ -187,8 +187,15 @@ namespace kllm { return; } - KaiRequest km_req; - auto rs = parse_oai_completions_request(request->body().to_string(), context, km_req); + CompletionBuilder builder; + auto rs = OaiParser::parse_completions_request(request->body().to_string(), builder); + if (!rs.ok()) { + auto err = format_aoi_error_response(rs); + res_error(response, err); + } + KaiRequest km_req = builder.final(); + KaiCompletions completion(context); + rs = completion.verify(km_req); if (!rs.ok()) { auto err = format_aoi_error_response(rs); res_error(response, err); @@ -196,7 +203,6 @@ namespace kllm { //// if (!km_req.slot_params().stream()) { KaiResponse km_res; - KaiCompletions completion(context); completion.completions(km_req, km_res); nlohmann::ordered_json arr = nlohmann::ordered_json::array();; for (auto &it: km_res.completions()) { @@ -209,7 +215,6 @@ namespace kllm { } else { KaiResponse km_res; - KaiCompletions completion(context); auto writer = response->get_chunk_streamer(); completion.completions_stream(km_req, km_res, [&writer, this](const CompletionsResult &result) -> bool { nlohmann::ordered_json data; @@ -238,15 +243,24 @@ namespace kllm { return; } - KaiRequest km_req; - auto rs = parse_oai_chat_completions_request(nlohmann::ordered_json::parse(request->body().to_string()), context, km_req); + ChatCompletionBuilder builder; + + auto rs = OaiParser::parse_chat_completions_request(nlohmann::ordered_json::parse(request->body().to_string()), builder); + if (!rs.ok()) { + auto err = format_aoi_error_response(rs); + res_error(response, err); + return; + } + KaiRequest km_req = builder.final(); + KaiChatCompletions chat(context); + rs = chat.verify(km_req); if (!rs.ok()) { auto err = format_aoi_error_response(rs); res_error(response, err); + return; } if(!km_req.slot_params().stream()) { krpc::ClosureGuard done_guard(done); - KaiChatCompletions chat(context); KaiResponse km_response; chat.chat(km_req, km_response); if(km_response.status().code() != 0) { @@ -258,10 +272,9 @@ namespace kllm { res_ok(response, body); } else { KaiResponse km_res; - KaiChatCompletions completion(context); auto writer = response->get_chunk_streamer(); done->Run(); - completion.chat_stream(km_req, km_res, [&](const CompletionsResult &result) -> bool { + chat.chat_stream(km_req, km_res, [&](const CompletionsResult &result) -> bool { auto data = format_partial_response_oaicompat(result, context); if(data.empty()) { return true; @@ -288,39 +301,21 @@ namespace kllm { } void InfillProcessor::process(const krpc::RestfulRequest *request, krpc::RestfulResponse *response) { - if (context->params.embedding) { - res_error(response, format_aoi_error_response( - "This server does not support completions. Start it without `--embeddings`", - ERROR_TYPE_NOT_SUPPORTED)); - return; - } - // check model compatibility - std::string err; - if (llama_token_fim_pre(context->km_model.model) == LLAMA_TOKEN_NULL) { - err += "prefix token is missing. "; - } - if (llama_token_fim_suf(context->km_model.model) == LLAMA_TOKEN_NULL) { - err += "suffix token is missing. "; - } - if (llama_token_fim_mid(context->km_model.model) == LLAMA_TOKEN_NULL) { - err += "middle token is missing. "; - } - if (!err.empty()) { - res_error(response, - format_aoi_error_response(turbo::str_format("Infill is not supported by this model: %s", err.c_str()), - ERROR_TYPE_NOT_SUPPORTED)); + InfillBuilder builder; + auto rs = OaiParser::parse_infill_request(request->body().to_string(), builder); + if(!rs.ok()) { + res_error(response, format_aoi_error_response(rs)); return; } - KaiRequest km_req; - auto rs = parse_oai_infill_request(request->body().to_string(), context,km_req); + KaiRequest km_req = builder.final(); + KaiInfill infill(context); + rs = infill.verify(km_req); if(!rs.ok()) { res_error(response, format_aoi_error_response(rs)); return; } - if (!km_req.slot_params().stream()) { KaiResponse km_res; - KaiInfill infill(context); infill.infill(km_req, km_res); nlohmann::ordered_json arr = nlohmann::ordered_json::array();; for (auto &it: km_res.completions()) { @@ -333,7 +328,6 @@ namespace kllm { } else { KaiResponse km_res; - KaiInfill infill(context); auto writer = response->get_chunk_streamer(); infill.infill_stream(km_req, km_res, [&](const CompletionsResult &result) -> bool { nlohmann::ordered_json data; @@ -356,10 +350,15 @@ namespace kllm { void EmbeddingProcessor::process(const krpc::RestfulRequest *request, krpc::RestfulResponse *response) { KaiEmbeddings embd(context); - KaiRequest km_req; + EmbeddingBuilder builder; KaiResponse km_res; - km_req.set_query_type(QUERY_EMBEDDING); - auto rs = parse_oai_embedding_request(request->body().to_string(), context, km_req); + auto rs = OaiParser::parse_embedding_request(request->body().to_string(), builder); + if (!rs.ok()) { + res_error(response, format_aoi_error_response(rs)); + return; + } + KaiRequest km_req = builder.final(); + rs = embd.verify(km_req); if (!rs.ok()) { res_error(response, format_aoi_error_response(rs)); return; @@ -389,16 +388,20 @@ namespace kllm { return; } - KaiRequest km_req; KaiResponse km_res; - auto rs = parse_oai_rerank_request(request->body().to_string(), context, km_req); + RerankBuilder builder; + auto rs = OaiParser::parse_rerank_request(request->body().to_string(), builder); if (!rs.ok()) { res_error(response, format_aoi_error_response(rs)); return; } - KaiRerank ranker(context); - + KaiRequest km_req = builder.final(); + rs = ranker.verify(km_req); + if (!rs.ok()) { + res_error(response, format_aoi_error_response(rs)); + return; + } ranker.rank(km_req, km_res); nlohmann::ordered_json json; rs = format_aoi_rerank_response(km_res, json); @@ -415,15 +418,21 @@ namespace kllm { } void TokenizeProcessor::process(const krpc::RestfulRequest *request, krpc::RestfulResponse *response) { - KaiRequest km_req; KaiResponse km_res; - auto rs = parse_oai_tokenize_request(request->body().to_string(), context, km_req); + TokenizeBuilder builder; + auto rs = OaiParser::parse_oai_tokenize_request(request->body().to_string(), builder); if (!rs.ok()) { res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); return; } - km_req.set_query_type(QUERY_TOKENIZE); + + KaiRequest km_req = builder.final(); KaiTokenize tokenize(context); + rs = tokenize.verify(km_req); + if (!rs.ok()) { + res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); + return; + } tokenize.tokenize(km_req, km_res); if (km_res.status().code() != 0) { res_error(response, format_aoi_error_response(km_res.status().errmsg(), ERROR_TYPE_INVALID_REQUEST)); @@ -443,15 +452,22 @@ namespace kllm { } void DetokenizeProcessor::process(const krpc::RestfulRequest *request, krpc::RestfulResponse *response) { - KaiRequest km_req; KaiResponse km_res; - auto rs = parse_oai_detokenize_request(response->body().to_string(), context, km_req); + DetokenizeBuilder builder; + auto rs = OaiParser::parse_detokenize_request(response->body().to_string(), builder); if (!rs.ok()) { res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); return; } - km_req.set_query_type(QUERY_DETOKENIZE); + + KaiRequest km_req = builder.final(); + KaiTokenize tokenize(context); + rs = tokenize.verify(km_req); + if (!rs.ok()) { + res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); + return; + } tokenize.detokenize(km_req, km_res); if (km_res.status().code() != 0) { res_error(response, format_aoi_error_response(km_res.status().errmsg(), ERROR_TYPE_INVALID_REQUEST)); @@ -474,6 +490,11 @@ namespace kllm { KaiRequest km_req; KaiResponse km_res; KaiLora lora(context); + auto rs = lora.verify(km_req); + if (!rs.ok()) { + res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); + return; + } lora.lora(km_req, km_res); if (km_res.status().code() != 0) { res_error(response, format_aoi_error_response(km_res.status().errmsg(), ERROR_TYPE_INVALID_REQUEST)); @@ -489,15 +510,21 @@ namespace kllm { } void LoraApplyProcessor::process(const krpc::RestfulRequest *request, krpc::RestfulResponse *response) { - KaiRequest km_req; KaiResponse km_res; - auto rs = parse_oai_lora_request(request->body().to_string(), context, km_req); + LoraBuilder builder; + auto rs = OaiParser::parse_lora_request(request->body().to_string(), builder); if (!rs.ok()) { res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); return; } KaiLora lora(context); + KaiRequest km_req = builder.final(); + rs = lora.verify(km_req); + if (!rs.ok()) { + res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); + return; + } lora.lora_apply(km_req, km_res); if (km_res.status().code() != 0) { res_error(response, format_aoi_error_response(km_res.status().errmsg(), ERROR_TYPE_INVALID_REQUEST)); @@ -564,34 +591,32 @@ namespace kllm { KaiRequest km_req; KaiResponse km_res; + SlotsSaveRequestBuilder builder; auto action_query = request->uri().GetQuery("action"); - std::string action = action_query ? *action_query : ""; - if (action == "save") { - km_req.set_query_type(QUERY_SLOTS_SAVE); - } else if (action == "restore") { - km_req.set_query_type(QUERY_SLOTS_RESTORE); - } else if (action == "erase") { - km_req.set_query_type(QUERY_SLOTS_ERASE); - } else { + if(!action_query) { + res_error(response, format_aoi_error_response("no action", ERROR_TYPE_INVALID_REQUEST)); + return; + } + auto rs = builder.set_action(*action_query); + if(!rs.ok()) { res_error(response, format_aoi_error_response("Invalid action", ERROR_TYPE_INVALID_REQUEST)); + return; } km_req.mutable_slots_task()->set_id_slot(id_slot); - auto rs = parse_oai_slots_request(request->body().to_string(), context, km_req); + rs = OaiParser::parse_slots_request(request->body().to_string(), builder); if (!rs.ok()) { auto err = format_aoi_error_response(rs); res_error(response, err); + return; } + km_req = builder.final(); KaiSlots slots(context); - if (action == "save") { - km_req.set_query_type(QUERY_SLOTS_SAVE); - slots.save(km_req, km_res); - } else if (action == "restore") { - km_req.set_query_type(QUERY_SLOTS_RESTORE); - slots.save(km_req, km_res); - } else if (action == "erase") { - km_req.set_query_type(QUERY_SLOTS_ERASE); - slots.erase(km_req, km_res); + rs = slots.verify(km_req); + if (!rs.ok()) { + auto err = format_aoi_error_response(rs); + res_error(response, err); } + slots.action(km_req, km_res); if(km_res.status().code() != 0) { auto err = format_aoi_error_response(km_res); res_error(response, err); diff --git a/kllm/openai/oai_processor.h b/kllm/openai/oai_processor.h index 784ec41..2f60336 100644 --- a/kllm/openai/oai_processor.h +++ b/kllm/openai/oai_processor.h @@ -33,7 +33,7 @@ namespace kllm { struct MetricProcessor : public krpc::RestfulProcessor { - MetricProcessor(KMContext *c) : context(c) {} + explicit MetricProcessor(KMContext *c) : context(c) {} ~MetricProcessor() override = default; @@ -46,7 +46,7 @@ namespace kllm { struct PropsProcessor : public krpc::RestfulProcessor { - PropsProcessor(KMContext *c) : context(c) {} + explicit PropsProcessor(KMContext *c) : context(c) {} ~PropsProcessor() override = default; @@ -63,7 +63,7 @@ namespace kllm { struct ModelsProcessor : public krpc::RestfulProcessor { - ModelsProcessor(KMContext *c) : context(c) {} + explicit ModelsProcessor(KMContext *c) : context(c) {} ~ModelsProcessor() override = default; @@ -76,7 +76,7 @@ namespace kllm { struct CompletionsProcessor : public krpc::RestfulProcessor { - CompletionsProcessor(KMContext *c) : context(c) {} + explicit CompletionsProcessor(KMContext *c) : context(c) {} ~CompletionsProcessor() override = default; @@ -89,7 +89,7 @@ namespace kllm { struct ChatCompletionsProcessor : public krpc::RestfulProcessor { - ChatCompletionsProcessor(KMContext *c) : context(c) {} + explicit ChatCompletionsProcessor(KMContext *c) : context(c) {} ~ChatCompletionsProcessor() override = default; @@ -110,7 +110,7 @@ namespace kllm { struct InfillProcessor : public krpc::RestfulProcessor { - InfillProcessor(KMContext *c) : context(c) {} + explicit InfillProcessor(KMContext *c) : context(c) {} ~InfillProcessor() override = default; @@ -123,7 +123,7 @@ namespace kllm { struct EmbeddingProcessor : public krpc::RestfulProcessor { - EmbeddingProcessor(KMContext *c) : context(c) {} + explicit EmbeddingProcessor(KMContext *c) : context(c) {} ~EmbeddingProcessor() override = default; @@ -136,7 +136,7 @@ namespace kllm { struct RerankProcessor : public krpc::RestfulProcessor { - RerankProcessor(KMContext *c) : context(c) {} + explicit RerankProcessor(KMContext *c) : context(c) {} ~RerankProcessor() override = default; @@ -149,7 +149,7 @@ namespace kllm { struct TokenizeProcessor : public krpc::RestfulProcessor { - TokenizeProcessor(KMContext *c) : context(c) {} + explicit TokenizeProcessor(KMContext *c) : context(c) {} ~TokenizeProcessor() override = default; @@ -162,7 +162,7 @@ namespace kllm { struct DetokenizeProcessor : public krpc::RestfulProcessor { - DetokenizeProcessor(KMContext *c) : context(c) {} + explicit DetokenizeProcessor(KMContext *c) : context(c) {} ~DetokenizeProcessor() override = default; @@ -175,7 +175,7 @@ namespace kllm { struct LoraListProcessor : public krpc::RestfulProcessor { - LoraListProcessor(KMContext *c) : context(c) {} + explicit LoraListProcessor(KMContext *c) : context(c) {} ~LoraListProcessor() override = default; @@ -188,7 +188,7 @@ namespace kllm { struct LoraApplyProcessor : public krpc::RestfulProcessor { - LoraApplyProcessor(KMContext *c) : context(c) {} + explicit LoraApplyProcessor(KMContext *c) : context(c) {} ~LoraApplyProcessor() override = default; @@ -209,12 +209,11 @@ namespace kllm { bool is_wildcards() const override; - KMContext *context; }; struct SlotsProcessor : public krpc::RestfulProcessor { - SlotsProcessor(KMContext *c) : context(c) {} + explicit SlotsProcessor(KMContext *c) : context(c) {} ~SlotsProcessor() override = default; diff --git a/kllm/openai/request.cc b/kllm/openai/parser.cc similarity index 38% rename from kllm/openai/request.cc rename to kllm/openai/parser.cc index b22fc86..d743f1f 100644 --- a/kllm/openai/request.cc +++ b/kllm/openai/parser.cc @@ -16,260 +16,159 @@ // // -#include +#include #include +#include #include -#include namespace kllm { - namespace internal { - turbo::Result - format_chat(const KMContext *ctx, const std::vector &messages) { - std::vector chat; - - for (size_t i = 0; i < messages.size(); ++i) { - const auto &curr_msg = messages[i]; - - std::string role = json_value(curr_msg, "role", std::string("")); - - std::string content; - if (curr_msg.contains("content")) { - if (curr_msg["content"].is_string()) { - content = curr_msg["content"].get(); - } else if (curr_msg["content"].is_array()) { - for (const auto &part: curr_msg["content"]) { - if (part.contains("text")) { - content += "\n" + part["text"].get(); - } - } - } else { - return turbo::invalid_argument_error( - "Invalid 'content' type (ref: https://github.com/ggerganov/llama.cpp/issues/8367)"); - } - } else { - return turbo::invalid_argument_error( - "Missing 'content' (ref: https://github.com/ggerganov/llama.cpp/issues/8367)"); - } - chat.push_back({role, content}); + turbo::Status + OaiParser::parse_slot_params(const nlohmann::ordered_json &data, DefaultRequestBuilder &req) { + static SlotParams default_params; + req.set_stream(json_value(data, "stream", false)); + req.set_cache_prompt(json_value(data, "cache_prompt", false)); + req.set_n_predict( + json_value(data, "n_predict", json_value(data, "max_tokens", -1))); + req.set_n_indent(json_value(data, "n_indent", default_params.n_indent())); + req.set_n_keep(json_value(data, "n_keep", default_params.n_keep())); + req.set_n_discard(json_value(data, "n_discard", default_params.n_discard())); + // TODO: implement + req.set_t_max_prompt_ms( + json_value(data, "t_max_prompt_ms", default_params.t_max_prompt_ms())); + req.set_t_max_predict_ms( + json_value(data, "t_max_predict_ms", default_params.t_max_predict_ms())); + + const auto &stop = data.find("stop"); + if (stop != data.end() && stop->is_array()) { + for (const auto &word: *stop) { + if (!word.empty()) { + req.append_anti_prompt(stop->get()); + } } - - const auto formatted_chat = ctx->chat_apply_template(ctx->params.chat_template, chat, true); - VLOG(300) << turbo::str_format("formatted_chat: '%s'\n", formatted_chat.c_str()); - - return formatted_chat; + } else if (stop != data.end() && stop->is_string()) { + req.append_anti_prompt(stop->get()); } + return turbo::OkStatus(); + } - static turbo::Status parse_aoi_input_request(const nlohmann::ordered_json &obj, KaiPrompts *request) { - auto icnt = obj.count("input"); - auto ccnt = obj.count("content"); + turbo::Status + OaiParser::parse_sampler_params(const nlohmann::ordered_json &data, DefaultRequestBuilder &req) { + static SampleParams default_sparams; + req.set_top_k(json_value(data, "top_k", default_sparams.top_k())); + req.set_top_p(json_value(data, "top_p", default_sparams.top_p())); + req.set_min_p(json_value(data, "min_p", default_sparams.min_p())); + req.set_xtc_probability( + json_value(data, "xtc_probability", default_sparams.xtc_probability())); + req.set_xtc_threshold(json_value(data, "xtc_threshold", default_sparams.xtc_threshold())); + req.set_typ_p(json_value(data, "typical_p", default_sparams.typ_p())); + req.set_temp(json_value(data, "temperature", default_sparams.temp())); + req.set_dynatemp_range( + json_value(data, "dynatemp_range", default_sparams.dynatemp_range())); + req.set_dynatemp_exponent( + json_value(data, "dynatemp_exponent", default_sparams.dynatemp_exponent())); + req.set_penalty_last_n( + json_value(data, "repeat_last_n", default_sparams.penalty_last_n())); + req.set_penalty_repeat( + json_value(data, "repeat_penalty", default_sparams.penalty_repeat())); + req.set_penalty_freq( + json_value(data, "frequency_penalty", default_sparams.penalty_freq())); + req.set_penalty_present( + json_value(data, "presence_penalty", default_sparams.penalty_present())); + req.set_dry_multiplier( + json_value(data, "dry_multiplier", default_sparams.dry_multiplier())); + req.set_dry_base(json_value(data, "dry_base", default_sparams.dry_base())); + req.set_dry_allowed_length( + json_value(data, "dry_allowed_length", default_sparams.dry_allowed_length())); + req.set_dry_penalty_last_n( + json_value(data, "dry_penalty_last_n", default_sparams.dry_penalty_last_n())); + req.set_mirostat(json_value(data, "mirostat", default_sparams.mirostat())); + req.set_mirostat_tau(json_value(data, "mirostat_tau", default_sparams.mirostat_tau())); + req.set_mirostat_eta(json_value(data, "mirostat_eta", default_sparams.mirostat_eta())); + req.set_penalize_nl(json_value(data, "penalize_nl", default_sparams.penalize_nl())); + req.set_seed(json_value(data, "seed", default_sparams.seed())); + req.set_n_probs(json_value(data, "n_probs", default_sparams.n_probs())); + req.set_min_keep(json_value(data, "min_keep", default_sparams.min_keep())); + if (data.contains("dry_sequence_breakers")) { + auto vl = json_value(data, "dry_sequence_breakers", + std::vector()); + req.append_dry_sequence_breakers(vl); + } + if (data.contains("json_schema") && !data.at("json_schema").is_null() && + data.contains("grammar") && !data.at("grammar").is_null()) { + return turbo::invalid_argument_error( + "Either \"json_schema\" or \"grammar\" can be specified, but not both"); + } + if (data.contains("json_schema") && !data.at("json_schema").is_null()) { + auto obj = json_value(data, "json_schema", nlohmann::ordered_json::object()); + req.set_grammar(json_schema_to_grammar(obj)); + } else { + req.set_grammar(json_value(data, "grammar", default_sparams.grammar())); + } - if (icnt != 0 && ccnt != 0) { - return turbo::invalid_argument_error("both \"input\" and \"content\" are provided, must the of these"); - } - nlohmann::ordered_json ipt; - if (icnt != 0) { - ipt = obj.at("input"); - } else { - ipt = obj.at("content"); - } - if (ipt.is_string()) { - PromptsValue value; - value.set_string_value(ipt.get()); - *request->mutable_values()->Add() = std::move(value); - return turbo::OkStatus(); + const auto &samplers = data.find("samplers"); + if (samplers != data.end() && samplers->is_array()) { + std::vector sampler_names; + for (const auto &name: *samplers) { + if (name.is_string()) { + sampler_names.push_back(name.get()); + } } - - if (ipt.is_number()) { - PromptsValue value; - value.set_number_value(ipt.get()); - *request->mutable_values()->Add() = std::move(value); - return turbo::OkStatus(); + auto r = common_sampler_types_from_names(sampler_names, false); + for(auto t : r) { + req.append_samplers(static_cast(t)); } - if (ipt.is_array()) { - for (auto &it: ipt) { - if (it.is_string()) { - PromptsValue value; - value.set_string_value(it.get()); - *request->mutable_values()->Add() = std::move(value); - continue; - } - if (it.is_number()) { - PromptsValue value; - value.set_number_value(it.get()); - *request->mutable_values()->Add() = std::move(value); - continue; - } - // list - PromptsValueList *list_value = request->mutable_values()->Add()->mutable_list_value(); - for (auto &sit: it) { - if (sit.is_string()) { - PromptsValue value; - value.set_string_value(sit.get()); - *list_value->mutable_values()->Add() = std::move(value); - continue; - } - if (sit.is_number()) { - PromptsValue value; - value.set_number_value(sit.get()); - *list_value->mutable_values()->Add() = std::move(value); - continue; - } - return turbo::invalid_argument_error("too more sub lists"); - } - } // for - } - return turbo::OkStatus(); } - turbo::Status - parse_oai_slot_params(const nlohmann::ordered_json &data, const KMContext *context, KaiRequest &req) { - static SlotParams default_params; - req.mutable_slot_params()->set_stream(json_value(data, "stream", false)); - req.mutable_slot_params()->set_cache_prompt(json_value(data, "cache_prompt", false)); - req.mutable_slot_params()->set_n_predict( - json_value(data, "n_predict", json_value(data, "max_tokens", context->params.n_predict))); - req.mutable_slot_params()->set_n_indent(json_value(data, "n_indent", default_params.n_indent())); - req.mutable_slot_params()->set_n_keep(json_value(data, "n_keep", default_params.n_keep())); - req.mutable_slot_params()->set_n_discard(json_value(data, "n_discard", default_params.n_discard())); - // TODO: implement - req.mutable_slot_params()->set_t_max_prompt_ms( - json_value(data, "t_max_prompt_ms", default_params.t_max_prompt_ms())); - req.mutable_slot_params()->set_t_max_predict_ms( - json_value(data, "t_max_predict_ms", default_params.t_max_predict_ms())); - - const auto &stop = data.find("stop"); - if (stop != data.end() && stop->is_array()) { - for (const auto &word: *stop) { - if (!word.empty()) { - *req.mutable_slot_params()->mutable_antiprompt()->Add() = word; - } - } - } else if (stop != data.end() && stop->is_string()) { - *req.mutable_slot_params()->mutable_antiprompt()->Add() = stop->get(); - } - return turbo::OkStatus(); + req.set_ignore_eos(false); + if (json_value(data, "ignore_eos", false)) { + req.set_ignore_eos(true); } - turbo::Status - parse_oai_sampler_params(const nlohmann::ordered_json &data, const KMContext *context, KaiRequest &req) { - auto &default_sparams = context->params.sparams; - auto sparam = req.mutable_sparam(); - sparam->set_top_k(json_value(data, "top_k", default_sparams.top_k)); - sparam->set_top_p(json_value(data, "top_p", default_sparams.top_p)); - sparam->set_min_p(json_value(data, "min_p", default_sparams.min_p)); - sparam->set_xtc_probability( - json_value(data, "xtc_probability", default_sparams.xtc_probability)); - sparam->set_xtc_threshold(json_value(data, "xtc_threshold", default_sparams.xtc_threshold)); - sparam->set_typ_p(json_value(data, "typical_p", default_sparams.typ_p)); - sparam->set_temp(json_value(data, "temperature", default_sparams.temp)); - sparam->set_dynatemp_range( - json_value(data, "dynatemp_range", default_sparams.dynatemp_range)); - sparam->set_dynatemp_exponent( - json_value(data, "dynatemp_exponent", default_sparams.dynatemp_exponent)); - sparam->set_penalty_last_n( - json_value(data, "repeat_last_n", default_sparams.penalty_last_n)); - sparam->set_penalty_repeat( - json_value(data, "repeat_penalty", default_sparams.penalty_repeat)); - sparam->set_penalty_freq( - json_value(data, "frequency_penalty", default_sparams.penalty_freq)); - sparam->set_penalty_present( - json_value(data, "presence_penalty", default_sparams.penalty_present)); - sparam->set_dry_multiplier( - json_value(data, "dry_multiplier", default_sparams.dry_multiplier)); - sparam->set_dry_base(json_value(data, "dry_base", default_sparams.dry_base)); - sparam->set_dry_allowed_length( - json_value(data, "dry_allowed_length", default_sparams.dry_allowed_length)); - sparam->set_dry_penalty_last_n( - json_value(data, "dry_penalty_last_n", default_sparams.dry_penalty_last_n)); - sparam->set_mirostat(json_value(data, "mirostat", default_sparams.mirostat)); - sparam->set_mirostat_tau(json_value(data, "mirostat_tau", default_sparams.mirostat_tau)); - sparam->set_mirostat_eta(json_value(data, "mirostat_eta", default_sparams.mirostat_eta)); - sparam->set_penalize_nl(json_value(data, "penalize_nl", default_sparams.penalize_nl)); - sparam->set_seed(json_value(data, "seed", default_sparams.seed)); - sparam->set_n_probs(json_value(data, "n_probs", default_sparams.n_probs)); - sparam->set_min_keep(json_value(data, "min_keep", default_sparams.min_keep)); - if (data.contains("dry_sequence_breakers")) { - auto vl = json_value(data, "dry_sequence_breakers", - std::vector()); - sparam->mutable_dry_sequence_breakers()->Assign(vl.begin(), vl.end()); - } - if (data.contains("json_schema") && !data.at("json_schema").is_null() && - data.contains("grammar") && !data.at("grammar").is_null()) { - return turbo::invalid_argument_error( - "Either \"json_schema\" or \"grammar\" can be specified, but not both"); - } - if (data.contains("json_schema") && !data.at("json_schema").is_null()) { - auto obj = json_value(data, "json_schema", nlohmann::ordered_json::object()); - sparam->set_grammar(json_schema_to_grammar(obj)); - } else { - sparam->set_grammar(json_value(data, "grammar", default_sparams.grammar)); - }; - - const auto &samplers = data.find("samplers"); - if (samplers != data.end() && samplers->is_array()) { - std::vector sampler_names; - for (const auto &name: *samplers) { - if (name.is_string()) { - sampler_names.push_back(name.get()); + const auto &logit_bias = data.find("logit_bias"); + if (logit_bias != data.end() && logit_bias->is_array()) { + for (const auto &el: *logit_bias) { + // TODO: we may want to throw errors here, in case "el" is incorrect + if (el.is_array() && el.size() == 2) { + float bias; + if (el[1].is_number()) { + bias = el[1].get(); + } else if (el[1].is_boolean() && !el[1].get()) { + bias = -INFINITY; + } else { + continue; } - } - auto r = common_sampler_types_from_names(sampler_names, false); - req.mutable_sparam()->mutable_samplers()->Assign(r.begin(), r.end()); - } else { - req.mutable_sparam()->mutable_samplers()->Assign(default_sparams.samplers.begin(), - default_sparams.samplers.end()); - } - - req.mutable_sparam()->set_ignore_eos(false); - if (json_value(data, "ignore_eos", false)) { - req.mutable_sparam()->set_ignore_eos(true); - } - const auto &logit_bias = data.find("logit_bias"); - if (logit_bias != data.end() && logit_bias->is_array()) { - for (const auto &el: *logit_bias) { - // TODO: we may want to throw errors here, in case "el" is incorrect - if (el.is_array() && el.size() == 2) { - float bias; - if (el[1].is_number()) { - bias = el[1].get(); - } else if (el[1].is_boolean() && !el[1].get()) { - bias = -INFINITY; - } else { - continue; - } - - if (el[0].is_number_integer()) { - llama_token tok = el[0].get(); - LogitBias lb; - lb.set_bias(bias); - lb.set_token(tok); - *req.mutable_sparam()->mutable_logit_bias()->Add() = lb; - } else if (el[0].is_string()) { - LogitBias lb; - lb.set_bias(bias); - lb.set_str_token(el[0].get()); - *req.mutable_sparam()->mutable_logit_bias()->Add() = lb; - } + if (el[0].is_number_integer()) { + llama_token tok = el[0].get(); + LogitBias lb; + lb.set_bias(bias); + lb.set_token(tok); + req.append_logit_bias(lb); + } else if (el[0].is_string()) { + LogitBias lb; + lb.set_bias(bias); + lb.set_str_token(el[0].get()); + req.append_logit_bias(lb); } } } + } - return turbo::OkStatus(); + return turbo::OkStatus(); + + } - } - } // namespace internal turbo::Status - parse_oai_base_request(const nlohmann::ordered_json &data, const KMContext *context, KaiRequest &req) { + OaiParser::parse_base_request(const nlohmann::ordered_json &data, DefaultRequestBuilder &req) { try { - auto rs = internal::parse_oai_slot_params(data, context, req); + auto rs = parse_slot_params(data, req); if (!rs.ok()) { return rs; } - rs = internal::parse_oai_sampler_params(data, context, req); + rs = parse_sampler_params(data, req); if (!rs.ok()) { return rs; } @@ -277,7 +176,7 @@ namespace kllm { req.set_id_slot(json_value(data, "id_slot", -1)); req.set_oaicompat(json_value(data, "__oaicompat", false)); req.set_fail_on_no_slot(json_value(data, "fail_on_no_slot", false)); - req.set_model(json_value(data, "model", context->params.model_alias)); + req.set_model(json_value(data, "model", std::string())); return turbo::OkStatus(); } catch (std::exception &e) { return turbo::unknown_error(e.what()); @@ -285,137 +184,250 @@ namespace kllm { } turbo::Status - parse_oai_tokenize_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req) { - auto rs = parse_oai_base_request(json, context, req); + OaiParser::parse_tokenize_request(const nlohmann::ordered_json &json, TokenizeBuilder &req) { + auto rs = parse_base_request(json, req); if (!rs.ok()) { return rs; } - rs = internal::parse_aoi_input_request(json, req.mutable_tokenize_request()->mutable_prompts()); - if (!rs.ok()) { - return rs; + + auto icnt = json.count("input"); + auto ccnt = json.count("content"); + + if (icnt != 0 && ccnt != 0) { + return turbo::invalid_argument_error("both \"input\" and \"content\" are provided, must the of these"); + } + + nlohmann::ordered_json ipt; + if (icnt != 0) { + ipt = json.at("input"); + } else { + ipt = json.at("content"); } - req.mutable_tokenize_request()->set_add_special(json_value(json, "add_special", false)); - req.mutable_tokenize_request()->set_with_pieces(json_value(json, "with_pieces", false)); - if (req.tokenize_request().prompts().values().empty()) { - return turbo::invalid_argument_error("\"input\" or \"content\"must be provided"); + if (ipt.is_string()) { + req.add_input(ipt.get()); + return turbo::OkStatus(); } + + if (ipt.is_number()) { + req.add_input(ipt.get()); + return turbo::OkStatus(); + } + + if (ipt.is_array()) { + for (auto &it: ipt) { + if (it.is_string()) { + req.add_input(it.get()); + continue; + } + if (it.is_number()) { + req.add_input(it.get()); + continue; + } + // list + std::vector> token_list; + for (auto &sit: it) { + if (sit.is_string()) { + std::variant value; + value = sit.get(); + token_list.push_back(value); + continue; + } + if (sit.is_number()) { + std::variant value; + value = sit.get(); + token_list.push_back(value); + continue; + } + return turbo::invalid_argument_error("too more sub lists"); + } + req.add_input(token_list); + } // for + } + + req.set_add_special(json_value(json, "add_special", false)); + req.set_with_pieces(json_value(json, "with_pieces", false)); return turbo::OkStatus(); } - turbo::Status parse_oai_tokenize_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_oai_tokenize_request(const std::string &content, TokenizeBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_tokenize_request(body, context, req); + return parse_tokenize_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } turbo::Status - parse_oai_detokenize_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req) { - auto rs = parse_oai_base_request(json, context, req); + OaiParser::parse_detokenize_request(const nlohmann::ordered_json &json, DetokenizeBuilder &req) { + auto rs = parse_base_request(json, req); if (!rs.ok()) { return rs; } if (json.count("tokens") != 0) { const std::vector tokens = json.at("tokens"); - req.mutable_detokenize()->mutable_tokens()->Assign(tokens.begin(), tokens.end()); + req.add_token(tokens); } return turbo::OkStatus(); } - turbo::Status parse_oai_detokenize_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_detokenize_request(const std::string &content, DetokenizeBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_detokenize_request(body, context, req); + return parse_detokenize_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } turbo::Status - parse_oai_embedding_request(const nlohmann::ordered_json &body, const KMContext *context, KaiRequest &req) { - auto rs = parse_oai_base_request(body, context, req); + OaiParser::parse_embedding_request(const nlohmann::ordered_json &body, EmbeddingBuilder &req) { + auto rs = parse_base_request(body, req); if (!rs.ok()) { return rs; } - rs = internal::parse_aoi_input_request(body, req.mutable_embedding_request()->mutable_prompts()); - if (!rs.ok()) { - return rs; + auto icnt = body.count("input"); + auto ccnt = body.count("content"); + + if (icnt != 0 && ccnt != 0) { + return turbo::invalid_argument_error("both \"input\" and \"content\" are provided, must the of these"); } - if (req.tokenize_request().prompts().values().empty()) { - return turbo::invalid_argument_error("\"input\" or \"content\"must be provided"); + nlohmann::ordered_json ipt; + if (icnt != 0) { + ipt = body.at("input"); + } else { + ipt = body.at("content"); + } + + if (ipt.is_string()) { + req.add_input(ipt.get()); + return turbo::OkStatus(); + } + + if (ipt.is_number()) { + req.add_input(ipt.get()); + return turbo::OkStatus(); + } + + if (ipt.is_array()) { + for (auto &it: ipt) { + if (it.is_string()) { + req.add_input(it.get()); + continue; + } + if (it.is_number()) { + req.add_input(it.get()); + continue; + } + // list + std::vector> token_list; + for (auto &sit: it) { + if (sit.is_string()) { + std::variant value; + value = sit.get(); + token_list.push_back(value); + continue; + } + if (sit.is_number()) { + std::variant value; + value = sit.get(); + token_list.push_back(value); + continue; + } + return turbo::invalid_argument_error("too more sub lists"); + } + req.add_input(token_list); + } // for } return turbo::OkStatus(); } - turbo::Status parse_oai_embedding_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_embedding_request(const std::string &content, EmbeddingBuilder &req) { try { LOG(INFO) << "content: " << content; const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_embedding_request(body, context, req); + return parse_embedding_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } turbo::Status - parse_oai_completions_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req) { - auto rs = parse_oai_base_request(json, context, req); + OaiParser::parse_completions_request(const nlohmann::ordered_json &json, CompletionBuilder &req) { + auto rs = parse_base_request(json, req); if (!rs.ok()) { return rs; } if (json.count("prompt")) { try { auto s = json.at("prompt").get(); - req.mutable_completion_request()->set_prompts(s); + req.set_prompt(s); } catch (const std::exception &e) { return turbo::invalid_argument_error("prompt must string"); } } - if (req.completion_request().prompts().empty()) { - return turbo::invalid_argument_error("\"prompt\" must be provided"); - } return turbo::OkStatus(); } - turbo::Status parse_oai_completions_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_completions_request(const std::string &content, CompletionBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_completions_request(body, context, req); + return parse_completions_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } turbo::Status - parse_oai_chat_completions_request(const nlohmann::ordered_json &body, const KMContext *context, KaiRequest &req) { - auto rs = parse_oai_base_request(body, context, req); + OaiParser::parse_chat_completions_request(const nlohmann::ordered_json &body, ChatCompletionBuilder &req) { + auto rs = parse_base_request(body, req); if (!rs.ok()) { return rs; } if (!body.contains("messages")) { return turbo::invalid_argument_error("miss messages field"); } - auto rst = internal::format_chat(context, body.at("messages")); - if (!rst.ok()) { - return rst.status(); + + auto messages = body.at("messages"); + //std::vector chat; + + for (const auto &curr_msg : messages) { + std::string role = json_value(curr_msg, "role", std::string("")); + + std::string content; + if (curr_msg.contains("content")) { + if (curr_msg["content"].is_string()) { + content = curr_msg["content"].get(); + } else if (curr_msg["content"].is_array()) { + for (const auto &part: curr_msg["content"]) { + if (part.contains("text")) { + content += "\n" + part["text"].get(); + } + } + } else { + return turbo::invalid_argument_error( + "Invalid 'content' type (ref: https://github.com/ggerganov/llama.cpp/issues/8367)"); + } + } else { + return turbo::invalid_argument_error( + "Missing 'content' (ref: https://github.com/ggerganov/llama.cpp/issues/8367)"); + } + req.append_content(role, content); } - req.mutable_chat_completion_request()->set_message(rst.value()); if (body.contains("response_format")) { nlohmann::ordered_json response_format = json_value(body, "response_format", nlohmann::ordered_json::object()); std::string response_type = json_value(response_format, "type", std::string()); if (response_type == "json_object") { - req.mutable_sparam()->set_grammar(json_schema_to_grammar( + req.set_grammar(json_schema_to_grammar( json_value(response_format, "schema", nlohmann::ordered_json::object()))); } else if (response_type == "json_schema") { nlohmann::ordered_json json_schema = json_value(response_format, "json_schema", nlohmann::ordered_json::object()); - req.mutable_sparam()->set_grammar( + req.set_grammar( json_schema_to_grammar(json_value(json_schema, "schema", nlohmann::ordered_json::object()))); } else if (!response_type.empty() && response_type != "text") { return turbo::invalid_argument_error( @@ -446,16 +458,16 @@ namespace kllm { } turbo::Status - parse_oai_chat_completions_request(const std::string &content, const KMContext *context, KaiRequest &req) { + OaiParser::parse_chat_completions_request(const std::string &content, ChatCompletionBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_chat_completions_request(body, context, req); + return parse_chat_completions_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } - turbo::Status parse_mix_value_list(const nlohmann::ordered_json &body, PromptsValue &value) { + static turbo::Status parse_mix_value_list(const nlohmann::ordered_json &body, PromptsValue &value) { for (auto &it: body) { if (it.is_string()) { PromptsValue v; @@ -473,9 +485,9 @@ namespace kllm { } turbo::Status - parse_oai_infill_request(const nlohmann::ordered_json &data, const KMContext *context, KaiRequest &req) { + OaiParser::parse_infill_request(const nlohmann::ordered_json &data, InfillBuilder &req) { auto body = data; - auto rs = parse_oai_base_request(body, context, req); + auto rs = parse_base_request(body, req); if (!rs.ok()) { return rs; } @@ -496,26 +508,25 @@ namespace kllm { if (!rs.ok()) { return turbo::invalid_argument_error("parse input prefix error"); } - *req.mutable_infill_request()->mutable_input_prefix() = input_prefix; + req.append_input_prefix(input_prefix); PromptsValue input_suffix; rs = parse_mix_value_list(body.at("input_suffix"), input_suffix); if (!rs.ok()) { return turbo::invalid_argument_error("parse input suffix error"); } - *req.mutable_infill_request()->mutable_input_suffix() = input_suffix; + req.append_input_suffix(input_suffix); if (!body.contains("input_extra")) { body["input_extra"] = nlohmann::ordered_json::array(); } try { for (auto &chunk: body) { - ExtraPair ep; const std::string text = json_value(chunk, "text", std::string()); const std::string filename = json_value(chunk, "filename", std::string("tmp")); - ep.set_filename(filename); - ep.set_text(text); - *req.mutable_infill_request()->mutable_input_extra()->Add() = std::move(ep); + if(!text.empty()) { + req.append_extra(filename, text); + } } } catch (const std::exception &e) { return turbo::data_loss_error(e.what()); @@ -524,30 +535,26 @@ namespace kllm { if (data.count("prompt")) { try { auto s = data.at("prompt").get(); - req.mutable_infill_request()->set_prompts(s); + req.set_prompt(s); } catch (const std::exception &e) { return turbo::invalid_argument_error("prompt must string"); } } - if (req.completion_request().prompts().empty()) { - return turbo::invalid_argument_error("\"prompt\" must be provided"); - } - return turbo::OkStatus(); } - turbo::Status parse_oai_infill_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_infill_request(const std::string &content, InfillBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_infill_request(body, context, req); + return parse_infill_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } turbo::Status - parse_oai_lora_request(const nlohmann::ordered_json &body, const KMContext *context, KaiRequest &req) { - auto rs = parse_oai_base_request(body, context, req); + OaiParser::parse_lora_request(const nlohmann::ordered_json &body, LoraBuilder &req) { + auto rs = parse_base_request(body, req); if (!rs.ok()) { return rs; } @@ -555,30 +562,31 @@ namespace kllm { const std::vector loras = body.at("loras"); std::string empty; for (const auto &it: loras) { - int id = it.at("id"); - float scale = it.at("scale"); - std::string path = json_value(it, "path", empty); - LoraInfo *info = req.mutable_lora_request()->mutable_lora_infos()->Add(); - info->set_id(id); - info->set_scale(scale); - info->set_path(path); + try { + int id = it.at("id"); + float scale = it.at("scale"); + std::string path = json_value(it, "path", empty); + req.append_lora(id, scale, path); + } catch (const std::exception &e) { + return turbo::invalid_argument_error("invalid lora request: %s", e.what()); + } } } return turbo::OkStatus(); } - turbo::Status parse_oai_lora_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_lora_request(const std::string &content, LoraBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_lora_request(body, context, req); + return parse_lora_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } turbo::Status - parse_oai_rerank_request(const nlohmann::ordered_json &body, const KMContext *context, KaiRequest &req) { - auto rs = parse_oai_base_request(body, context, req); + OaiParser::parse_rerank_request(const nlohmann::ordered_json &body, RerankBuilder &req) { + auto rs = parse_base_request(body, req); if (!rs.ok()) { return rs; } @@ -598,35 +606,24 @@ namespace kllm { return turbo::invalid_argument_error("\"documents\" must be a non-empty string array"); } - req.mutable_rerank_request()->set_query(query.get()); + req.set_query(query.get()); for (auto &it: documents) { - *req.mutable_rerank_request()->mutable_docs()->Add() = it; + req.append_docs(it); } - // merge to PromptsValue - auto prompt = req.mutable_rerank_request()->mutable_prompts()->mutable_values()->Add(); - PromptsValue pquery; - pquery.set_string_value(req.rerank_request().query()); - *prompt->mutable_list_value()->mutable_values()->Add() = std::move(pquery); - for (auto &it: documents) { - PromptsValue doc; - doc.set_string_value(it); - *prompt->mutable_list_value()->mutable_values()->Add() = std::move(doc); - } - return turbo::OkStatus(); } - turbo::Status parse_oai_rerank_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_rerank_request(const std::string &content, RerankBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_rerank_request(body, context, req); + return parse_rerank_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } } turbo::Status - parse_oai_slots_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req) { + OaiParser::parse_slots_request(const nlohmann::ordered_json &json, SlotsSaveRequestBuilder &req) { if (req.query_type() == QUERY_SLOTS_ERASE) { return turbo::OkStatus(); } @@ -637,16 +634,14 @@ namespace kllm { if (!fs_validate_filename(filename)) { return turbo::invalid_argument_error("Invalid filename"); } - std::string filepath = context->params.slot_save_path + filename; - req.mutable_slots_task()->set_filename(filename); - req.mutable_slots_task()->set_filepath(filepath); + req.set_filename(filename); return turbo::OkStatus(); } - turbo::Status parse_oai_slots_request(const std::string &content, const KMContext *context, KaiRequest &req) { + turbo::Status OaiParser::parse_slots_request(const std::string &content, SlotsSaveRequestBuilder &req) { try { const nlohmann::ordered_json body = nlohmann::ordered_json::parse(content); - return parse_oai_slots_request(body, context, req); + return parse_slots_request(body, req); } catch (std::exception &e) { return turbo::data_loss_error("bad json: %s", e.what()); } diff --git a/kllm/openai/request.h b/kllm/openai/parser.h similarity index 30% rename from kllm/openai/request.h rename to kllm/openai/parser.h index df31989..2ba2c35 100644 --- a/kllm/openai/request.h +++ b/kllm/openai/parser.h @@ -18,68 +18,78 @@ #pragma once -#include +#include #include #include #include -#include namespace kllm { namespace internal { turbo::Status - parse_oai_slot_params(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + parse_oai_slot_params(const nlohmann::ordered_json &json, KaiRequest &req); turbo::Status - parse_oai_sampler_params(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + parse_oai_sampler_params(const nlohmann::ordered_json &json, KaiRequest &req); } - turbo::Status parse_oai_base_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); - turbo::Status - parse_oai_tokenize_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + //turbo::Status parse_oai_base_request(const nlohmann::ordered_json &json, TokenizeBuilder &req); - turbo::Status parse_oai_tokenize_request(const std::string &content, const KMContext *context, KaiRequest &req); + class OaiParser { + public: - turbo::Status - parse_oai_detokenize_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status parse_base_request(const nlohmann::ordered_json &json, DefaultRequestBuilder &req); - turbo::Status parse_oai_detokenize_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_tokenize_request(const nlohmann::ordered_json &json, TokenizeBuilder &req); + static turbo::Status parse_oai_tokenize_request(const std::string &content, TokenizeBuilder &req); - turbo::Status - parse_oai_embedding_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_detokenize_request(const nlohmann::ordered_json &json, DetokenizeBuilder &req); - turbo::Status parse_oai_embedding_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status parse_detokenize_request(const std::string &content, DetokenizeBuilder &req); - turbo::Status - parse_oai_completions_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_embedding_request(const nlohmann::ordered_json &json, EmbeddingBuilder &req); - turbo::Status parse_oai_completions_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status parse_embedding_request(const std::string &content, EmbeddingBuilder &req); - turbo::Status - parse_oai_chat_completions_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_slots_request(const nlohmann::ordered_json &json, SlotsSaveRequestBuilder &req); - turbo::Status - parse_oai_chat_completions_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status parse_slots_request(const std::string &content, SlotsSaveRequestBuilder &req); - turbo::Status - parse_oai_infill_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_rerank_request(const nlohmann::ordered_json &json, RerankBuilder &req); - turbo::Status parse_oai_infill_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status parse_rerank_request(const std::string &content, RerankBuilder &req); - turbo::Status - parse_oai_lora_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_lora_request(const nlohmann::ordered_json &json, LoraBuilder &req); - turbo::Status parse_oai_lora_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status parse_lora_request(const std::string &content, LoraBuilder &req); - turbo::Status - parse_oai_rerank_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_infill_request(const nlohmann::ordered_json &json, InfillBuilder &req); - turbo::Status parse_oai_rerank_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status parse_infill_request(const std::string &content, InfillBuilder &req); - turbo::Status - parse_oai_slots_request(const nlohmann::ordered_json &json, const KMContext *context, KaiRequest &req); + static turbo::Status + parse_completions_request(const nlohmann::ordered_json &json, CompletionBuilder &req); - turbo::Status parse_oai_slots_request(const std::string &content, const KMContext *context, KaiRequest &req); + static turbo::Status parse_completions_request(const std::string &content, CompletionBuilder &req); + static turbo::Status + parse_chat_completions_request(const nlohmann::ordered_json &json, ChatCompletionBuilder &req); + + static turbo::Status + parse_chat_completions_request(const std::string &content,ChatCompletionBuilder &req); + + private: + static turbo::Status parse_slot_params(const nlohmann::ordered_json &json, DefaultRequestBuilder &req); + + static turbo::Status + parse_sampler_params(const nlohmann::ordered_json &json, DefaultRequestBuilder &req); + }; } // namespace kllm diff --git a/kllm/proto/builder.cc b/kllm/proto/builder.cc new file mode 100644 index 0000000..fb69c2b --- /dev/null +++ b/kllm/proto/builder.cc @@ -0,0 +1,95 @@ +// Copyright (C) 2024 Kumo inc. +// Author: Jeff.li lijippy@163.com +// All rights reserved. +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published +// by the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// + +#include + +namespace kllm { + + DefaultRequestBuilder &DefaultRequestBuilder::setup(const SlotParams* sp, const SampleParams*sample) { + static SlotParams default_params; + if(sp != nullptr) { + setup_slot_params(sp); + } + setup_sampler_params(sample); + return *this; + } + + void DefaultRequestBuilder::setup_slot_params(const SlotParams* sp) { + *_request.mutable_slot_params() = *sp; + } + + void DefaultRequestBuilder::setup_sampler_params(const SampleParams* sp) { + if(sp != nullptr) { + *_request.mutable_sparam() = *sp; + return; + } + /* + static InternalSamplerParams isp; + _request.mutable_sparam()->mutable_dry_sequence_breakers()->Assign(isp.dry_sequence_breakers.begin(), isp.dry_sequence_breakers.end()); + _request.mutable_sparam()->mutable_samplers()->Assign(isp.samplers.begin(), isp.samplers.end()); + */ + } + + std::vector + common_sampler_types_from_names(const std::vector &names, bool allow_alt_names) { + std::unordered_map sampler_canonical_name_map{ + {"dry", COMMON_SAMPLER_TYPE_DRY}, + {"top_k", COMMON_SAMPLER_TYPE_TOP_K}, + {"top_p", COMMON_SAMPLER_TYPE_TOP_P}, + {"typ_p", COMMON_SAMPLER_TYPE_TYPICAL_P}, + {"min_p", COMMON_SAMPLER_TYPE_MIN_P}, + {"temperature", COMMON_SAMPLER_TYPE_TEMPERATURE}, + {"xtc", COMMON_SAMPLER_TYPE_XTC}, + {"infill", COMMON_SAMPLER_TYPE_INFILL}, + }; + + // since samplers names are written multiple ways + // make it ready for both system names and input names + std::unordered_map sampler_alt_name_map{ + {"top-k", COMMON_SAMPLER_TYPE_TOP_K}, + {"top-p", COMMON_SAMPLER_TYPE_TOP_P}, + {"nucleus", COMMON_SAMPLER_TYPE_TOP_P}, + {"typical-p", COMMON_SAMPLER_TYPE_TYPICAL_P}, + {"typical", COMMON_SAMPLER_TYPE_TYPICAL_P}, + {"typ-p", COMMON_SAMPLER_TYPE_TYPICAL_P}, + {"typ", COMMON_SAMPLER_TYPE_TYPICAL_P}, + {"min-p", COMMON_SAMPLER_TYPE_MIN_P}, + {"temp", COMMON_SAMPLER_TYPE_TEMPERATURE}, + }; + + std::vector samplers; + samplers.reserve(names.size()); + + for (const auto &name: names) { + auto sampler = sampler_canonical_name_map.find(name); + if (sampler != sampler_canonical_name_map.end()) { + samplers.push_back(sampler->second); + } else { + if (allow_alt_names) { + sampler = sampler_alt_name_map.find(name); + if (sampler != sampler_alt_name_map.end()) { + samplers.push_back(sampler->second); + } + } + } + } + + return samplers; + } + +} // namespace kllm diff --git a/kllm/proto/builder.h b/kllm/proto/builder.h new file mode 100644 index 0000000..edd6e7e --- /dev/null +++ b/kllm/proto/builder.h @@ -0,0 +1,944 @@ +// Copyright (C) 2024 Kumo inc. +// Author: Jeff.li lijippy@163.com +// All rights reserved. +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published +// by the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// + +#pragma once + +#include +#include +#include +#include + +namespace kllm { + + class DefaultRequestBuilder { + public: + DefaultRequestBuilder() = default; + + virtual ~DefaultRequestBuilder() = default; + + KaiRequest final() { + merge(); + return std::move(_request); + } + + DefaultRequestBuilder &setup(const SlotParams *sp = nullptr, const SampleParams *sample = nullptr); + + /// slot params + DefaultRequestBuilder &enable_stream() { + _request.mutable_slot_params()->set_stream(true); + return *this; + } + + DefaultRequestBuilder &disable_stream() { + _request.mutable_slot_params()->set_stream(false); + return *this; + } + + DefaultRequestBuilder &set_stream(bool flag) { + _request.mutable_slot_params()->set_stream(flag); + return *this; + } + + DefaultRequestBuilder &base() { + return *this; + } + + bool stream() const { + return _request.slot_params().stream(); + } + + DefaultRequestBuilder &enable_cache_prompt() { + _request.mutable_slot_params()->set_cache_prompt(true); + return *this; + } + + DefaultRequestBuilder &disable_cache_prompt() { + _request.mutable_slot_params()->set_cache_prompt(false); + return *this; + } + + DefaultRequestBuilder &enable_fail_on_no_slot() { + _request.set_fail_on_no_slot(true); + return *this; + } + + DefaultRequestBuilder &disable_fail_on_no_slot() { + _request.set_fail_on_no_slot(false); + return *this; + } + + DefaultRequestBuilder &set_fail_on_no_slot(bool flag) { + _request.set_fail_on_no_slot(flag); + return *this; + } + + DefaultRequestBuilder &enable_oaicompat() { + _request.set_oaicompat(true); + return *this; + } + + DefaultRequestBuilder &disable_oaicompat() { + _request.set_oaicompat(false); + return *this; + } + + DefaultRequestBuilder &set_oaicompat(bool flag) { + _request.set_oaicompat(flag); + return *this; + } + + DefaultRequestBuilder &set_id_slot(int32_t id) { + _request.set_id_slot(id); + return *this; + } + + DefaultRequestBuilder &set_model(const std::string &m) { + _request.set_model(m); + return *this; + } + + DefaultRequestBuilder &set_cache_prompt(bool flag) { + _request.mutable_slot_params()->set_cache_prompt(flag); + return *this; + } + + bool cache_prompt() const { + return _request.slot_params().cache_prompt(); + } + + DefaultRequestBuilder &set_n_keep(int32_t n) { + _request.mutable_slot_params()->set_n_keep(n); + return *this; + } + + int32_t n_keep() const { + return _request.slot_params().n_keep(); + } + + DefaultRequestBuilder &set_n_discard(int32_t n) { + _request.mutable_slot_params()->set_n_discard(n); + return *this; + } + + int32_t n_discard() const { + return _request.slot_params().n_discard(); + } + + DefaultRequestBuilder &set_n_predict(int32_t n) { + _request.mutable_slot_params()->set_n_predict(n); + return *this; + } + + int32_t n_predict() const { + return _request.slot_params().n_predict(); + } + + DefaultRequestBuilder &set_n_indent(int32_t n) { + _request.mutable_slot_params()->set_n_indent(n); + return *this; + } + + int32_t n_indent() const { + return _request.slot_params().n_indent(); + } + + DefaultRequestBuilder &set_t_max_prompt_ms(int32_t n) { + _request.mutable_slot_params()->set_t_max_prompt_ms(n); + return *this; + } + + int32_t t_max_prompt_ms() const { + return _request.slot_params().t_max_prompt_ms(); + } + + DefaultRequestBuilder &set_t_max_predict_ms(int32_t n) { + _request.mutable_slot_params()->set_t_max_predict_ms(n); + return *this; + } + + int32_t t_max_predict_ms() const { + return _request.slot_params().t_max_predict_ms(); + } + + + DefaultRequestBuilder &clear_anti_prompt() { + _request.mutable_slot_params()->mutable_antiprompt()->Clear(); + return *this; + } + + const ::google::protobuf::RepeatedPtrField &anti_prompt() const { + return _request.slot_params().antiprompt(); + } + + DefaultRequestBuilder &append_anti_prompt(const std::string &word) { + *_request.mutable_slot_params()->mutable_antiprompt()->Add() = word; + return *this; + } + + DefaultRequestBuilder &append_anti_prompt(const std::vector &words) { + for (auto &it: words) { + *_request.mutable_slot_params()->mutable_antiprompt()->Add() = it; + } + return *this; + } + + DefaultRequestBuilder &set_anti_prompt(const std::string &word) { + _request.mutable_slot_params()->mutable_antiprompt()->Clear(); + *_request.mutable_slot_params()->mutable_antiprompt()->Add() = word; + return *this; + } + + DefaultRequestBuilder &set_anti_prompt(const std::vector &words) { + _request.mutable_slot_params()->mutable_antiprompt()->Assign( + words.begin(), + words.end() + ); + return *this; + } + + /// samplers + + DefaultRequestBuilder &set_seed(int32_t n) { + _request.mutable_sparam()->set_seed(n); + return *this; + } + + DefaultRequestBuilder &set_n_prev(int32_t n) { + _request.mutable_sparam()->set_n_prev(n); + return *this; + } + + DefaultRequestBuilder &set_n_probs(int32_t n) { + _request.mutable_sparam()->set_n_probs(n); + return *this; + } + + DefaultRequestBuilder &set_min_keep(int32_t n) { + _request.mutable_sparam()->set_min_keep(n); + return *this; + } + + DefaultRequestBuilder &set_top_k(int32_t n) { + _request.mutable_sparam()->set_top_k(n); + return *this; + } + + DefaultRequestBuilder &set_top_p(float n) { + _request.mutable_sparam()->set_top_p(n); + return *this; + } + + DefaultRequestBuilder &set_min_p(float n) { + _request.mutable_sparam()->set_min_p(n); + return *this; + } + + DefaultRequestBuilder &set_xtc_probability(float n) { + _request.mutable_sparam()->set_xtc_probability(n); + return *this; + } + + DefaultRequestBuilder &set_xtc_threshold(float n) { + _request.mutable_sparam()->set_xtc_threshold(n); + return *this; + } + + DefaultRequestBuilder &set_typ_p(float n) { + _request.mutable_sparam()->set_typ_p(n); + return *this; + } + + DefaultRequestBuilder &set_temp(float n) { + _request.mutable_sparam()->set_temp(n); + return *this; + } + + DefaultRequestBuilder &set_dynatemp_range(float n) { + _request.mutable_sparam()->set_dynatemp_range(n); + return *this; + } + + DefaultRequestBuilder &set_dynatemp_exponent(float n) { + _request.mutable_sparam()->set_dynatemp_exponent(n); + return *this; + } + + DefaultRequestBuilder &set_penalty_last_n(int32_t n) { + _request.mutable_sparam()->set_penalty_last_n(n); + return *this; + } + + DefaultRequestBuilder &set_penalty_repeat(float n) { + _request.mutable_sparam()->set_penalty_repeat(n); + return *this; + } + + DefaultRequestBuilder &set_penalty_freq(float n) { + _request.mutable_sparam()->set_penalty_freq(n); + return *this; + } + + DefaultRequestBuilder &set_penalty_present(float n) { + _request.mutable_sparam()->set_penalty_present(n); + return *this; + } + + DefaultRequestBuilder &set_dry_multiplier(float n) { + _request.mutable_sparam()->set_dry_multiplier(n); + return *this; + } + + DefaultRequestBuilder &set_dry_base(float n) { + _request.mutable_sparam()->set_dry_base(n); + return *this; + } + + DefaultRequestBuilder &set_dry_allowed_length(int32_t n) { + _request.mutable_sparam()->set_dry_allowed_length(n); + return *this; + } + + DefaultRequestBuilder &set_dry_penalty_last_n(int32_t n) { + _request.mutable_sparam()->set_dry_penalty_last_n(n); + return *this; + } + + DefaultRequestBuilder &set_mirostat(int32_t n) { + _request.mutable_sparam()->set_mirostat(n); + return *this; + } + + DefaultRequestBuilder &set_mirostat_tau(float n) { + _request.mutable_sparam()->set_mirostat_tau(n); + return *this; + } + + DefaultRequestBuilder &set_mirostat_eta(float n) { + _request.mutable_sparam()->set_mirostat_eta(n); + return *this; + } + + DefaultRequestBuilder &set_penalize_nl(bool flag) { + _request.mutable_sparam()->set_penalize_nl(flag); + return *this; + } + + DefaultRequestBuilder &enable_penalize_nl() { + _request.mutable_sparam()->set_penalize_nl(true); + return *this; + } + + DefaultRequestBuilder &disable_penalize_nl() { + _request.mutable_sparam()->set_penalize_nl(false); + return *this; + } + + DefaultRequestBuilder &set_ignore_eos(bool flag) { + _request.mutable_sparam()->set_ignore_eos(flag); + return *this; + } + + DefaultRequestBuilder &enable_ignore_eos() { + _request.mutable_sparam()->set_ignore_eos(true); + return *this; + } + + DefaultRequestBuilder &disable_ignore_eos() { + _request.mutable_sparam()->set_ignore_eos(false); + return *this; + } + + DefaultRequestBuilder &set_no_perf(bool flag) { + _request.mutable_sparam()->set_no_perf(flag); + return *this; + } + + DefaultRequestBuilder &enable_no_perf() { + _request.mutable_sparam()->set_no_perf(true); + return *this; + } + + DefaultRequestBuilder &disable_no_perf() { + _request.mutable_sparam()->set_no_perf(false); + return *this; + } + + DefaultRequestBuilder &clear_dry_sequence_breakers() { + _request.mutable_sparam()->mutable_dry_sequence_breakers()->Clear(); + return *this; + } + + const ::google::protobuf::RepeatedPtrField &dry_sequence_breakers() const { + return _request.sparam().dry_sequence_breakers(); + } + + DefaultRequestBuilder &append_dry_sequence_breakers(const std::string &word) { + *_request.mutable_sparam()->mutable_dry_sequence_breakers()->Add() = word; + return *this; + } + + DefaultRequestBuilder &append_dry_sequence_breakers(const std::vector &words) { + for (auto &it: words) { + *_request.mutable_sparam()->mutable_dry_sequence_breakers()->Add() = it; + } + return *this; + } + + DefaultRequestBuilder &set_dry_sequence_breakers(const std::string &word) { + _request.mutable_sparam()->mutable_dry_sequence_breakers()->Clear(); + *_request.mutable_sparam()->mutable_dry_sequence_breakers()->Add() = word; + return *this; + } + + DefaultRequestBuilder &set_dry_sequence_breakers(const std::vector &words) { + _request.mutable_sparam()->mutable_dry_sequence_breakers()->Assign( + words.begin(), + words.end() + ); + return *this; + } + + DefaultRequestBuilder &clear_samplers() { + _request.mutable_sparam()->mutable_samplers()->Clear(); + return *this; + } + + const ::google::protobuf::RepeatedField &samplers() const { + return _request.sparam().samplers(); + } + + DefaultRequestBuilder &append_samplers(int word) { + *_request.mutable_sparam()->mutable_samplers()->Add() = word; + return *this; + } + + DefaultRequestBuilder &append_samplers(const std::vector &words) { + for (auto &it: words) { + *_request.mutable_sparam()->mutable_samplers()->Add() = it; + } + return *this; + } + + DefaultRequestBuilder &set_samplers(const int32_t &word) { + _request.mutable_sparam()->mutable_samplers()->Clear(); + *_request.mutable_sparam()->mutable_samplers()->Add() = word; + return *this; + } + + DefaultRequestBuilder &set_samplers(const std::vector &words) { + _request.mutable_sparam()->mutable_samplers()->Assign( + words.begin(), + words.end() + ); + return *this; + } + + DefaultRequestBuilder &set_grammar(const std::string &word) { + _request.mutable_sparam()->set_grammar(word); + return *this; + } + + DefaultRequestBuilder &clear_logit_bias() { + _request.mutable_sparam()->mutable_logit_bias()->Clear(); + return *this; + } + + DefaultRequestBuilder &append_logit_bias(const LogitBias &lb) { + *_request.mutable_sparam()->mutable_logit_bias()->Add() = lb; + return *this; + } + + DefaultRequestBuilder &append_logit_bias(const std::vector &lbs) { + _request.mutable_sparam()->mutable_logit_bias()->Assign(lbs.begin(), lbs.end()); + return *this; + } + + DefaultRequestBuilder &append_logit_bias(int32_t token, float bias, const std::string &str = "") { + LogitBias lb; + lb.set_str_token(str); + lb.set_bias(bias); + lb.set_token(token); + *_request.mutable_sparam()->mutable_logit_bias()->Add() = lb; + return *this; + } + + const ::google::protobuf::RepeatedPtrField &logit_bias() const { + return _request.sparam().logit_bias(); + } + + const SampleParams &sample_params() const { + return _request.sparam(); + } + + SampleParams *mutable_sample_params() { + return _request.mutable_sparam(); + } + + const SlotParams &slot_params() const { + return _request.slot_params(); + } + + SlotParams *slot_params() { + return _request.mutable_slot_params(); + } + + private: + void setup_slot_params(const SlotParams *sp); + + void setup_sampler_params(const SampleParams *sp); + + virtual void merge() { + + } + + protected: + KaiRequest _request; + + }; + + class TokenizeBuilder : public DefaultRequestBuilder { + public: + TokenizeBuilder() { + setup(); + _request.set_query_type(QUERY_TOKENIZE); + } + + TokenizeBuilder &add_input(int32_t token) { + _request.mutable_tokenize_request()->mutable_prompts()->mutable_values()->Add()->set_number_value(token); + return *this; + } + + TokenizeBuilder &add_input(const std::vector &tokens) { + PromptsValue value; + for (auto it: tokens) { + PromptsValue t; + t.set_number_value(it); + *value.mutable_list_value()->mutable_values()->Add() = std::move(t); + } + *_request.mutable_tokenize_request()->mutable_prompts()->mutable_values()->Add() = value; + return *this; + } + + TokenizeBuilder &add_input(const std::vector &tokens) { + PromptsValue value; + for (auto &it: tokens) { + PromptsValue t; + t.set_string_value(it); + *value.mutable_list_value()->mutable_values()->Add() = std::move(t); + } + *_request.mutable_tokenize_request()->mutable_prompts()->mutable_values()->Add() = value; + return *this; + } + + TokenizeBuilder &add_input(const std::vector> &tokens) { + PromptsValue value; + for (auto &it: tokens) { + PromptsValue t; + const int *int_ptr = std::get_if(&it); + const std::string *str_ptr = std::get_if(&it); + if (int_ptr) { + t.set_number_value(*int_ptr); + } else { + t.set_string_value(*str_ptr); + } + *value.mutable_list_value()->mutable_values()->Add() = std::move(t); + } + *_request.mutable_tokenize_request()->mutable_prompts()->mutable_values()->Add() = value; + return *this; + } + + TokenizeBuilder &add_input(const std::string &piece) { + _request.mutable_tokenize_request()->mutable_prompts()->mutable_values()->Add()->set_string_value(piece); + return *this; + } + + TokenizeBuilder &clear_input() { + _request.mutable_tokenize_request()->mutable_prompts()->mutable_values()->Clear(); + return *this; + } + + + TokenizeBuilder &set_add_special(bool flag) { + _request.mutable_tokenize_request()->set_add_special(flag); + return *this; + } + + TokenizeBuilder &set_with_pieces(bool flag) { + _request.mutable_tokenize_request()->set_with_pieces(flag); + return *this; + } + }; + + class DetokenizeBuilder : public DefaultRequestBuilder { + public: + DetokenizeBuilder() { + setup(); + _request.set_query_type(QUERY_DETOKENIZE); + } + + DetokenizeBuilder &add_token(int32_t token) { + *_request.mutable_detokenize()->mutable_tokens()->Add() = token; + return *this; + } + + DetokenizeBuilder &clear_token() { + _request.mutable_detokenize()->mutable_tokens()->Clear(); + return *this; + } + + DetokenizeBuilder &add_token(const std::vector &tokens) { + _request.mutable_detokenize()->mutable_tokens()->Assign(tokens.begin(), tokens.end()); + return *this; + } + }; + + class SlotsSaveRequestBuilder : public DefaultRequestBuilder { + public: + SlotsSaveRequestBuilder() { + setup(); + } + + ~SlotsSaveRequestBuilder() override = default; + + turbo::Status set_action(const std::string &action) { + if (action == "save") { + _request.set_query_type(QUERY_SLOTS_SAVE); + } else if (action == "restore") { + _request.set_query_type(QUERY_SLOTS_RESTORE); + } else if (action == "erase") { + _request.set_query_type(QUERY_SLOTS_ERASE); + } else { + return turbo::invalid_argument_error("Invalid action"); + } + return turbo::OkStatus(); + } + + QueryType query_type() const { + return _request.query_type(); + } + + turbo::Status set_action(QueryType type) { + if(type == QUERY_SLOTS_SAVE || type == QUERY_SLOTS_RESTORE || type == QUERY_SLOTS_ERASE) { + _request.set_query_type(type); + return turbo::OkStatus(); + } + return turbo::invalid_argument_error("Invalid action"); + } + + SlotsSaveRequestBuilder &set_filename(const std::string &filename) { + _request.mutable_slots_task()->set_filename(filename); + return *this; + } + + SlotsSaveRequestBuilder &set_path(const std::string &filepath) { + _request.mutable_slots_task()->set_filepath(filepath); + return *this; + } + + SlotsSaveRequestBuilder &set_slot_id(int32_t slot_id) { + _request.mutable_slots_task()->set_id_slot(slot_id); + return *this; + } + + }; + + class RerankBuilder : public DefaultRequestBuilder { + public: + RerankBuilder() { + setup(); + _request.set_query_type(QUERY_TOKENIZE); + } + + RerankBuilder &set_query(const std::string &query) { + _request.mutable_rerank_request()->set_query(query); + return *this; + } + + RerankBuilder &clear_docs() { + _request.mutable_rerank_request()->mutable_docs()->Clear(); + return *this; + } + + RerankBuilder &append_docs(const std::string &doc) { + *_request.mutable_rerank_request()->mutable_docs()->Add() = doc; + return *this; + } + + RerankBuilder &append_docs(const std::vector &docs) { + for (auto &it: docs) { + *_request.mutable_rerank_request()->mutable_docs()->Add() = it; + } + return *this; + } + + RerankBuilder &set_docs(const std::vector &docs) { + _request.mutable_rerank_request()->mutable_docs()->Clear(); + for (auto &it: docs) { + *_request.mutable_rerank_request()->mutable_docs()->Add() = it; + } + return *this; + } + + private: + void merge() override { + auto prompt = _request.mutable_rerank_request()->mutable_prompts()->mutable_values()->Add(); + PromptsValue pquery; + pquery.set_string_value(_request.rerank_request().query()); + *prompt->mutable_list_value()->mutable_values()->Add() = std::move(pquery); + for (auto &it: _request.rerank_request().docs()) { + PromptsValue doc; + doc.set_string_value(it); + *prompt->mutable_list_value()->mutable_values()->Add() = std::move(doc); + } + } + }; + + class LoraBuilder : public DefaultRequestBuilder { + public: + LoraBuilder() { + setup(); + _request.set_query_type(QUERY_TOKENIZE); + } + + LoraBuilder &clear_lora() { + _request.mutable_lora_request()->mutable_lora_infos()->Clear(); + return *this; + } + + LoraBuilder &append_lora(int id, float scale, const std::string &path) { + auto l = _request.mutable_lora_request()->mutable_lora_infos()->Add(); + l->set_path(path); + l->set_scale(scale); + l->set_id(id); + return *this; + } + + + }; + + + class InfillBuilder : public DefaultRequestBuilder { + public: + InfillBuilder() { + setup(); + _request.set_query_type(QUERY_INFILL); + } + + InfillBuilder &set_prompt(const std::string &prompt) { + _request.mutable_infill_request()->set_prompts(prompt); + return *this; + } + + InfillBuilder &clear_input_prefix() { + _request.mutable_infill_request()->mutable_input_prefix()->Clear(); + return *this; + } + + InfillBuilder &append_input_prefix(const std::string &input) { + _request.mutable_infill_request()->mutable_input_prefix()->mutable_list_value()->mutable_values()->Add()->set_string_value( + input); + return *this; + } + + InfillBuilder &append_input_prefix(int32_t token) { + _request.mutable_infill_request()->mutable_input_prefix()->mutable_list_value()->mutable_values()->Add()->set_number_value( + token); + return *this; + } + + InfillBuilder &append_input_prefix(const PromptsValue &value) { + *_request.mutable_infill_request()->mutable_input_prefix()->mutable_list_value()->mutable_values()->Add() = value; + return *this; + } + + InfillBuilder &clear_input_suffix() { + _request.mutable_infill_request()->mutable_input_suffix()->Clear(); + return *this; + } + + InfillBuilder &append_input_suffix(const std::string &input) { + _request.mutable_infill_request()->mutable_input_suffix()->mutable_list_value()->mutable_values()->Add()->set_string_value( + input); + return *this; + } + + InfillBuilder &append_input_suffix(int32_t token) { + _request.mutable_infill_request()->mutable_input_suffix()->mutable_list_value()->mutable_values()->Add()->set_number_value( + token); + return *this; + } + + InfillBuilder &append_input_suffix(const PromptsValue &value) { + *_request.mutable_infill_request()->mutable_input_suffix()->mutable_list_value()->mutable_values()->Add() = value; + return *this; + } + + InfillBuilder &clear_extra() { + _request.mutable_infill_request()->mutable_input_extra(); + return *this; + } + + InfillBuilder &append_extra(const std::string &text, const std::string &filename) { + auto et = _request.mutable_infill_request()->mutable_input_extra()->Add(); + et->set_filename(filename); + et->set_text(text); + return *this; + } + + }; + + class EmbeddingBuilder : public DefaultRequestBuilder { + public: + EmbeddingBuilder() { + setup(); + _request.set_query_type(QUERY_EMBEDDING); + } + + EmbeddingBuilder &add_input(int32_t token) { + _request.mutable_embedding_request()->mutable_prompts()->mutable_values()->Add()->set_number_value(token); + return *this; + } + + EmbeddingBuilder &add_input(const std::vector &tokens) { + PromptsValue value; + for (auto it: tokens) { + PromptsValue t; + t.set_number_value(it); + *value.mutable_list_value()->mutable_values()->Add() = std::move(t); + } + *_request.mutable_embedding_request()->mutable_prompts()->mutable_values()->Add() = value; + return *this; + } + + EmbeddingBuilder &add_input(const std::vector &tokens) { + PromptsValue value; + for (auto &it: tokens) { + PromptsValue t; + t.set_string_value(it); + *value.mutable_list_value()->mutable_values()->Add() = std::move(t); + } + *_request.mutable_embedding_request()->mutable_prompts()->mutable_values()->Add() = value; + return *this; + } + + EmbeddingBuilder &add_input(const std::vector> &tokens) { + PromptsValue value; + for (auto &it: tokens) { + PromptsValue t; + const int *int_ptr = std::get_if(&it); + const std::string *str_ptr = std::get_if(&it); + if (int_ptr) { + t.set_number_value(*int_ptr); + } else { + t.set_string_value(*str_ptr); + } + *value.mutable_list_value()->mutable_values()->Add() = std::move(t); + } + *_request.mutable_embedding_request()->mutable_prompts()->mutable_values()->Add() = value; + return *this; + } + + EmbeddingBuilder &add_input(const std::string &piece) { + _request.mutable_embedding_request()->mutable_prompts()->mutable_values()->Add()->set_string_value(piece); + return *this; + } + + EmbeddingBuilder &clear_input() { + _request.mutable_embedding_request()->mutable_prompts()->mutable_values()->Clear(); + return *this; + } + }; + + class CompletionBuilder : public DefaultRequestBuilder { + public: + CompletionBuilder() { + setup(); + _request.set_query_type(QUERY_COMPLETION); + } + + CompletionBuilder &set_prompt(const std::string &promt) { + _request.mutable_completion_request()->set_prompts(promt); + return *this; + } + }; + + class ChatCompletionBuilder : public DefaultRequestBuilder { + public: + ChatCompletionBuilder() { + setup(); + _request.set_query_type(QUERY_COMPLETION); + } + + /// if use this api, no content* api is no effect + ChatCompletionBuilder&set_prompt(const std::string &promt) { + _request.mutable_chat_completion_request()->set_message(promt); + return *this; + } + + ChatCompletionBuilder&clear_content() { + _request.mutable_chat_completion_request()->mutable_segments()->Clear(); + return *this; + } + + ChatCompletionBuilder&append_content(const ChatMessage &seg) { + *_request.mutable_chat_completion_request()->mutable_segments()->Add() = seg; + return *this; + } + + ChatCompletionBuilder&append_content(const std::string& role, const std::string &content) { + auto seg = _request.mutable_chat_completion_request()->mutable_segments()->Add(); + seg->set_role(role); + seg->set_content(content); + return *this; + } + private: + void merge() override { + if(!_request.chat_completion_request().message().empty()) { + return; + } + auto &tem = _request.chat_completion_request().segments(); + turbo::flat_hash_map> seg; + for(auto &item: tem) { + auto it = seg.find(item.role()); + if(it != seg.end()) { + it->second.push_back(item.content()); + continue; + } + seg[item.role()] = {item.content()}; + } + _request.mutable_chat_completion_request()->mutable_segments()->Clear(); + for(auto &sit : seg) { + ChatMessage msg; + msg.set_role(sit.first); + if(sit.second.size() >1) { + std::string c; + for(auto &text : sit.second) { + c = c + text + "\n"; + } + msg.set_content(c); + } else { + msg.set_content(sit.second.front()); + } + *_request.mutable_chat_completion_request()->mutable_segments()->Add() = std::move(msg); + } + } + }; + + std::vector + common_sampler_types_from_names(const std::vector &names, bool allow_alt_names); +} // namespace kllm diff --git a/kllm/proto/interface.struct.proto b/kllm/proto/interface.struct.proto index dde6970..f5e4cf7 100644 --- a/kllm/proto/interface.struct.proto +++ b/kllm/proto/interface.struct.proto @@ -48,6 +48,10 @@ enum QueryType { QUERY_SLOTS_SAVE = 7; QUERY_SLOTS_RESTORE = 8; QUERY_SLOTS_ERASE = 9; + QUERY_INFILL = 10; + QUERY_COMPLETION = 11; + QUERY_CHAT_COMPLETION = 12; + QUERY_RERANK = 13; }; message ModelMeta { required LlamaVocabType vocab_type = 1; @@ -145,83 +149,83 @@ message LogitBias { message SampleParams { // the seed used to initialize llama_sampler - required int32 seed = 1; + required int32 seed = 1 [default = -1]; // number of previous tokens to remember - required int32 n_prev = 2; + required int32 n_prev = 2 [default = 64]; // if greater than 0, output the probabilities of top n_probs tokens. - required int32 n_probs = 3; + required int32 n_probs = 3 [default = 0]; // 0 = disabled, otherwise samplers should return at least min_keep tokens - required int32 min_keep = 4; + required int32 min_keep = 4 [default = 0]; // <= 0 to use vocab size - required int32 top_k = 5; + required int32 top_k = 5 [default = 40]; // 1.0 = disabled - required float top_p = 6; + required float top_p = 6 [default = 0.95]; // 0.0 = disabled - required float min_p = 7; + required float min_p = 7 [default = 0.05]; // 0.0 = disabled - required float xtc_probability = 8; + required float xtc_probability = 8 [default = 0.0]; // > 0.5 disables XTC - required float xtc_threshold = 9; + required float xtc_threshold = 9 [default = 0.10]; // typical_p, 1.0 = disabled - required float typ_p = 10; + required float typ_p = 10 [default = 1.0]; // <= 0.0 to sample greedily, 0.0 to not output probabilities - required float temp = 11; + required float temp = 11 [default = 0.80]; // 0.0 = disabled - required float dynatemp_range = 12; + required float dynatemp_range = 12 [default = 0.0]; // controls how entropy maps to temperature in dynamic temperature sampler - required float dynatemp_exponent = 13; + required float dynatemp_exponent = 13 [default = 1.0]; // last n tokens to penalize (0 = disable penalty, -1 = context size) - required int32 penalty_last_n = 14; + required int32 penalty_last_n = 14 [default = 64]; // 1.0 = disabled - required float penalty_repeat = 15; + required float penalty_repeat = 15 [default = 1.0]; // 0.0 = disabled - required float penalty_freq = 16; + required float penalty_freq = 16 [default = 0.0]; // 0.0 = disabled - required float penalty_present = 17; + required float penalty_present = 17 [default = 0.0]; // 0.0 = disabled; DRY repetition penalty for tokens extending repetition: - required float dry_multiplier = 18; + required float dry_multiplier = 18 [default = 0.0]; // 0.0 = disabled; multiplier * base ^ (length of sequence before token - allowed length) - required float dry_base = 19; + required float dry_base = 19 [default = 1.75]; // tokens extending repetitions beyond this receive penalty - required int32 dry_allowed_length = 20; + required int32 dry_allowed_length = 20 [default = 2]; // how many tokens to scan for repetitions (0 = disable penalty, -1 = context size) - required int32 dry_penalty_last_n = 21; + required int32 dry_penalty_last_n = 21 [default = -1]; // 0 = disabled, 1 = mirostat, 2 = mirostat 2.0 - required int32 mirostat = 22; + required int32 mirostat = 22 [default = 0]; // target entropy - required float mirostat_tau = 23; + required float mirostat_tau = 23 [default = 5.0]; // learning rate - required float mirostat_eta = 24; + required float mirostat_eta = 24 [default = 0.1]; // consider newlines as a repeatable token - required bool penalize_nl = 25; - required bool ignore_eos = 26; + required bool penalize_nl = 25 [default = false]; + required bool ignore_eos = 26 [default = false]; // disable performance metrics - required bool no_perf = 27; + required bool no_perf = 27 [default = false]; // default sequence breakers for DRY repeated string dry_sequence_breakers = 28; @@ -229,7 +233,7 @@ message SampleParams { repeated KaiSamplerType samplers = 29; // optional BNF-like grammar to constrain sampling - required string grammar = 30; + required string grammar = 30 [default = ""]; // logit biases to apply repeated LogitBias logit_bias = 31; @@ -352,8 +356,14 @@ message CompletionRequest { required string prompts = 1; } +message ChatMessage { + required string role = 1; + required string content = 2; +}; + message ChatCompletionRequest { required string message = 1; + repeated ChatMessage segments = 2; } message RerankRequest { @@ -371,8 +381,7 @@ message KaiRequest { optional string model = 5; optional bool fail_on_no_slot = 6; required int32 id_slot = 7 [default = -1]; - optional int32 n_probs = 8; // no use now - + required bool verified = 8 [default = false]; optional RerankRequest rerank_request = 9; optional TokenizeRequest tokenize_request = 10; diff --git a/kllm/tools/embedding/emb.cc b/kllm/tools/embedding/emb.cc index c2c224d..9c9f52b 100644 --- a/kllm/tools/embedding/emb.cc +++ b/kllm/tools/embedding/emb.cc @@ -19,8 +19,8 @@ #include #include #include -#include #include +#include namespace kllm { @@ -67,17 +67,18 @@ namespace kllm { KaiEmbeddings embd(&ServiceContext::ctx_server); - KaiRequest km_req; KaiResponse km_res; - km_req.set_query_type(QUERY_EMBEDDING); + EmbeddingBuilder builder; std::vector prompts = split_emb_lines(ServiceContext::params.prompt, ServiceContext::params.embd_sep); - kllm::PromptsValue string_list; for(auto &it : prompts) { - kllm::PromptsValue entity; - entity.set_string_value(it); - *string_list.mutable_list_value()->mutable_values()->Add() = std::move(entity); + builder.add_input(it); + } + KaiRequest km_req = builder.final(); + auto rs = embd.verify(km_req); + if(!rs.ok()) { + std::cerr<mutable_prompts()->mutable_values()->Add() = string_list; oai_context.start_context_async(); embd.embedding(km_req, km_res); // write JSON response @@ -85,7 +86,7 @@ namespace kllm { nlohmann::ordered_json json; int err_code; std::string errmsg; - format_aoi_embedding_response(km_res, json, errmsg, err_code); + kllm::format_aoi_embedding_response(km_res, json, errmsg, err_code); std::cout<> converter; -- Gitee From 1c393446abc8b206a0c48a1d0fa37310122cdab6 Mon Sep 17 00:00:00 2001 From: Jeff lothar Date: Mon, 2 Dec 2024 23:29:02 +0800 Subject: [PATCH 2/6] add cmd tokenize --- README.md | 5 ++ kllm/core/queue.h | 32 +++++---- kllm/kai/tokenize.cc | 4 ++ kllm/openai/oai_processor.cc | 4 +- kllm/openai/tokenize.cc | 8 ++- kllm/openai/tokenize.h | 4 +- kllm/proto/interface.struct.proto | 1 + kllm/tools/embedding/emb.cc | 12 ++-- kllm/tools/service_context.cc | 7 ++ kllm/tools/tokenize/CMakeLists.txt | 5 ++ kllm/tools/tokenize/tokenize.cc | 106 +++++++++++++++++++++++++++++ kllm/tools/tokenize/tokenize.h | 30 ++++++++ 12 files changed, 193 insertions(+), 25 deletions(-) create mode 100644 kllm/tools/tokenize/CMakeLists.txt create mode 100644 kllm/tools/tokenize/tokenize.cc create mode 100644 kllm/tools/tokenize/tokenize.h diff --git a/README.md b/README.md index fe40818..a5b6723 100644 --- a/README.md +++ b/README.md @@ -102,4 +102,9 @@ curl http://localhost:8080/v1/chat/completions \ ```shell ./kllm/kllm embedding --model /home/jeff/gitee/kumo-pub/temp/hf/qwen2.5-coder-1.5b-instruct-q5_0.gguf -p 我是宙斯 +``` + +### tokenize +```shell +./kllm/kllm tokenize --model /home/jeff/gitee/kumo-pub/temp/hf/qwen2.5-coder-1.5b-instruct-q5_0.gguf -p jieba分词 ``` \ No newline at end of file diff --git a/kllm/core/queue.h b/kllm/core/queue.h index 70dbcd1..501a57e 100644 --- a/kllm/core/queue.h +++ b/kllm/core/queue.h @@ -20,13 +20,14 @@ #include "kllm/core/metric.h" #include #include - +#include namespace kllm { struct server_queue { int id = 0; - bool running; + std::atomic running; + std::atomic notify_stop{false}; // queues std::deque queue_tasks; @@ -37,7 +38,7 @@ namespace kllm { // callback functions std::function callback_new_task; - std::function callback_update_slots; + std::function callback_update_slots; // Add a new task to the end of the queue int post(server_task task, bool front = false) { @@ -56,9 +57,9 @@ namespace kllm { } // multi-task version of post() - int post(std::vector & tasks, bool front = false) { + int post(std::vector &tasks, bool front = false) { std::unique_lock lock(mutex_tasks); - for (auto & task : tasks) { + for (auto &task: tasks) { if (task.id == -1) { task.id = id++; } @@ -110,9 +111,12 @@ namespace kllm { // end the start_loop routine void terminate() { - std::unique_lock lock(mutex_tasks); - running = false; - condition_tasks.notify_all(); + while (!notify_stop.load(std::memory_order_acquire)) { + running = false; + usleep(10000); + std::unique_lock lock(mutex_tasks); + condition_tasks.notify_all(); + } } /** @@ -127,7 +131,10 @@ namespace kllm { while (true) { QUE_DBG("%s", "processing new tasks\n"); - + if(!running.load(std::memory_order_acquire)) { + notify_stop = true; + return; + } while (true) { std::unique_lock lock(mutex_tasks); if (queue_tasks.empty()) { @@ -151,12 +158,13 @@ namespace kllm { { std::unique_lock lock(mutex_tasks); if (queue_tasks.empty()) { - if (!running) { + if (!running.load(std::memory_order_acquire)) { + notify_stop = true; QUE_DBG("%s", "terminate\n"); return; } - condition_tasks.wait(lock, [&]{ - return (!queue_tasks.empty() || !running); + condition_tasks.wait(lock, [&] { + return (!queue_tasks.empty() || !running.load(std::memory_order_acquire)); }); } } diff --git a/kllm/kai/tokenize.cc b/kllm/kai/tokenize.cc index 575b58a..dfdfc4f 100644 --- a/kllm/kai/tokenize.cc +++ b/kllm/kai/tokenize.cc @@ -34,6 +34,10 @@ namespace kllm { for(auto &l1 :tokens ) { for(auto &l2 : l1) { res.mutable_tokens()->Add(l2); + if(req.tokenize_request().with_pieces()) { + std::string piece = _context->token_to_piece(l2); + *res.mutable_tokens_piece()->Add() = std::move(piece); + } } } res.set_with_pieces(with_pieces); diff --git a/kllm/openai/oai_processor.cc b/kllm/openai/oai_processor.cc index 6180dd2..d85fc79 100644 --- a/kllm/openai/oai_processor.cc +++ b/kllm/openai/oai_processor.cc @@ -439,7 +439,7 @@ namespace kllm { return; } nlohmann::ordered_json obj; - rs = format_aoi_tokenize_response(km_res, context, obj); + rs = format_aoi_tokenize_response(km_res, obj); if (!rs.ok()) { res_error(response, format_aoi_error_response(rs.to_string(), ERROR_TYPE_INVALID_REQUEST)); return; @@ -474,7 +474,7 @@ namespace kllm { return; } nlohmann::ordered_json obj; - rs = format_aoi_detokenize_response(km_res, context, obj); + rs = format_aoi_detokenize_response(km_res, obj); if (!rs.ok()) { res_error(response, format_aoi_error_response(rs)); return; diff --git a/kllm/openai/tokenize.cc b/kllm/openai/tokenize.cc index 3816656..7a2e209 100644 --- a/kllm/openai/tokenize.cc +++ b/kllm/openai/tokenize.cc @@ -20,13 +20,14 @@ namespace kllm { - turbo::Status format_aoi_tokenize_response(const KaiResponse &response, KMContext *ctx, nlohmann::ordered_json &json_res) { + turbo::Status format_aoi_tokenize_response(const KaiResponse &response, nlohmann::ordered_json &json_res) { auto with_pieces = response.with_pieces(); auto &tokens = response.tokens(); auto obj = nlohmann::ordered_json::array(); if(with_pieces) { + int i = 0; for (const auto& token : tokens) { - std::string piece = ctx->token_to_piece(token); + std::string piece = response.tokens_piece(i); nlohmann::ordered_json piece_json; // Check if the piece is valid UTF-8 @@ -44,6 +45,7 @@ namespace kllm { {"id", token}, {"piece", piece_json} }); + ++i; } } else { obj = tokens; @@ -52,7 +54,7 @@ namespace kllm { return turbo::OkStatus(); } - turbo::Status format_aoi_detokenize_response(const KaiResponse &response, KMContext *ctx, nlohmann::ordered_json &obj) { + turbo::Status format_aoi_detokenize_response(const KaiResponse &response, nlohmann::ordered_json &obj) { obj = nlohmann::ordered_json{ {"content", response.content()} }; diff --git a/kllm/openai/tokenize.h b/kllm/openai/tokenize.h index 9d912b1..24b3766 100644 --- a/kllm/openai/tokenize.h +++ b/kllm/openai/tokenize.h @@ -26,8 +26,8 @@ namespace kllm { - turbo::Status format_aoi_tokenize_response(const KaiResponse &response, KMContext *ctx, nlohmann::ordered_json &obj); + turbo::Status format_aoi_tokenize_response(const KaiResponse &response, nlohmann::ordered_json &obj); - turbo::Status format_aoi_detokenize_response(const KaiResponse &response, KMContext *ctx, nlohmann::ordered_json &obj); + turbo::Status format_aoi_detokenize_response(const KaiResponse &response, nlohmann::ordered_json &obj); } // namespace kllm diff --git a/kllm/proto/interface.struct.proto b/kllm/proto/interface.struct.proto index f5e4cf7..935f779 100644 --- a/kllm/proto/interface.struct.proto +++ b/kllm/proto/interface.struct.proto @@ -454,4 +454,5 @@ message KaiResponse { repeated RerankResult reranks = 14; optional SlotsTaskResult slots_task = 15; optional string model_alias = 16; + repeated string tokens_piece = 17; } diff --git a/kllm/tools/embedding/emb.cc b/kllm/tools/embedding/emb.cc index 9c9f52b..851817f 100644 --- a/kllm/tools/embedding/emb.cc +++ b/kllm/tools/embedding/emb.cc @@ -15,7 +15,7 @@ // along with this program. If not, see . // -#include +#include #include #include #include @@ -40,12 +40,12 @@ namespace kllm { return lines; } - static ServiceContext oai_context; + static ServiceContext emb_context; static void run_embedding(); turbo::Status setup_embedding_cmd(turbo::cli::App *app) { auto emb_cmd = app->add_subcommand("embedding", "embedding the inputs"); - oai_context.params_context = ParamsContext::setup_app_context(emb_cmd,ServiceContext::params,LLAMA_EXAMPLE_EMBEDDING); + emb_context.params_context = ParamsContext::setup_app_context(emb_cmd,ServiceContext::params,LLAMA_EXAMPLE_EMBEDDING); turbo::Servlet::setup_log_option(emb_cmd); emb_cmd->callback(run_embedding); return turbo::OkStatus(); @@ -61,7 +61,7 @@ namespace kllm { // Generally you only need one Server. // load the model LOG_INF("%s: loading model\n", __func__); - oai_context.state.store(SERVER_STATE_READY); + emb_context.state.store(SERVER_STATE_READY); LOG_INF("%s: model loaded\n", __func__); @@ -79,7 +79,7 @@ namespace kllm { std::cerr< #include #include +#include #include namespace kllm { @@ -65,6 +66,7 @@ namespace kllm { if(ctx_runner) { ctx_server.queue_tasks.terminate(); ctx_runner->join(); + ctx_runner.reset(); } } @@ -83,6 +85,11 @@ namespace kllm { return rs; } + rs = setup_tokenize_cmd(app); + if(!rs.ok()) { + return rs; + } + //app->require_subcommand(1); app->parse_complete_callback([app](){ auto subs = app->get_subcommands(); diff --git a/kllm/tools/tokenize/CMakeLists.txt b/kllm/tools/tokenize/CMakeLists.txt new file mode 100644 index 0000000..b704dca --- /dev/null +++ b/kllm/tools/tokenize/CMakeLists.txt @@ -0,0 +1,5 @@ +set(TARGET llama-tokenize) +add_executable(${TARGET} tokenize.cpp) +install(TARGETS ${TARGET} RUNTIME) +target_link_libraries(${TARGET} PRIVATE common llama ${CMAKE_THREAD_LIBS_INIT}) +target_compile_features(${TARGET} PRIVATE cxx_std_11) diff --git a/kllm/tools/tokenize/tokenize.cc b/kllm/tools/tokenize/tokenize.cc new file mode 100644 index 0000000..0a283ec --- /dev/null +++ b/kllm/tools/tokenize/tokenize.cc @@ -0,0 +1,106 @@ +// Copyright (C) 2024 Kumo inc. +// Author: Jeff.li lijippy@163.com +// All rights reserved. +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published +// by the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// + +#include +#include +#include +#include +#include +#include + +namespace kllm { + + static std::vector split_emb_lines(const std::string &s, const std::string &separator = "\n") { + std::vector lines; + size_t start = 0; + size_t end = s.find(separator); + + while (end != std::string::npos) { + lines.push_back(s.substr(start, end - start)); + start = end + separator.length(); + end = s.find(separator, start); + } + + lines.push_back(s.substr(start)); // Add the last part + + return lines; + } + + static ServiceContext tk_context; + + static void run_tokenize(); + + turbo::Status setup_tokenize_cmd(turbo::cli::App *app) { + auto tk_cmd = app->add_subcommand("tokenize", "The tokenize program tokenizes a prompt using a given model,\n" + "and prints the resulting tokens to standard output.\n\n" + "It needs a model file, a prompt, and optionally other flags\n" + "to control the behavior of the tokenizer.\n\n"); + tk_context.params_context = ParamsContext::setup_app_context(tk_cmd, ServiceContext::params, + LLAMA_EXAMPLE_EMBEDDING); + turbo::Servlet::setup_log_option(tk_cmd); + tk_cmd->callback(run_tokenize); + return turbo::OkStatus(); + } + + void run_tokenize() { + + ServiceContext::params.embedding = true; + // For non-causal models, batch size must be equal to ubatch size + ServiceContext::params.n_ubatch = ServiceContext::params.n_batch; + + ServiceContext::call_after_parse(); + // Generally you only need one Server. + // load the model + LOG_INF("%s: loading model\n", __func__); + tk_context.state.store(SERVER_STATE_READY); + LOG_INF("%s: model loaded\n", __func__); + + + KaiTokenize tk(&ServiceContext::ctx_server); + KaiResponse km_res; + TokenizeBuilder builder; + std::vector prompts = split_emb_lines(ServiceContext::params.prompt, + ServiceContext::params.embd_sep); + for (auto &it: prompts) { + builder.add_input(it); + } + builder.set_with_pieces(true); + KaiRequest km_req = builder.final(); + auto rs = tk.verify(km_req); + if (!rs.ok()) { + std::cerr << rs.to_string() << std::endl; + return; + } + tk_context.start_context_async(); + tk.tokenize(km_req, km_res); + // write JSON response + if (km_res.status().code() == 0) { + nlohmann::ordered_json json; + rs = kllm::format_aoi_tokenize_response(km_res, json); + if(!rs.ok()) { + std::cerr<<"error: "<. +// + +#pragma once + +#include +#include +#include +#include + +namespace kllm { + + turbo::Status setup_tokenize_cmd(turbo::cli::App *app); + +} // namespace kllm + -- Gitee From 305d7c04a483bda3d34e8551803704217a440f6e Mon Sep 17 00:00:00 2001 From: Jeff lothar Date: Fri, 10 Jan 2025 02:47:15 +0800 Subject: [PATCH 3/6] fix readme --- README.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index a5b6723..39c7dfc 100644 --- a/README.md +++ b/README.md @@ -101,10 +101,10 @@ curl http://localhost:8080/v1/chat/completions \ ### embedding ```shell - ./kllm/kllm embedding --model /home/jeff/gitee/kumo-pub/temp/hf/qwen2.5-coder-1.5b-instruct-q5_0.gguf -p 我是宙斯 + ./kllm/kllm embedding --model qwen2.5-coder-1.5b-instruct-q5_0.gguf -p 我是宙斯 ``` ### tokenize ```shell -./kllm/kllm tokenize --model /home/jeff/gitee/kumo-pub/temp/hf/qwen2.5-coder-1.5b-instruct-q5_0.gguf -p jieba分词 +./kllm/kllm tokenize --model qwen2.5-coder-1.5b-instruct-q5_0.gguf -p jieba分词 ``` \ No newline at end of file -- Gitee From e17085fc91de7839cdced43c4351c01d99b6096f Mon Sep 17 00:00:00 2001 From: Jeff lothar Date: Fri, 10 Jan 2025 03:10:36 +0800 Subject: [PATCH 4/6] update --- cmake/kllm_deps.cmake | 5 ++- kllm/config/arg.cc | 2 +- kllm/core/params.cc | 2 +- kllm/openai/oai_processor.cc | 60 +++++++++++++++++------------------- kllm/openai/oai_processor.h | 4 +-- kllm/tools/kllm.cc | 2 +- kllm/tools/oai/oai.cc | 8 ++--- kmpkg-configuration.json | 2 +- kmpkg.json | 2 +- 9 files changed, 43 insertions(+), 44 deletions(-) diff --git a/cmake/kllm_deps.cmake b/cmake/kllm_deps.cmake index c518304..804a401 100644 --- a/cmake/kllm_deps.cmake +++ b/cmake/kllm_deps.cmake @@ -93,7 +93,10 @@ endif () ########################################################## set(KMCMAKE_DEPS_LINK turbo::turbo_static - krpc::krpc_static + tally::tally_static + alioth::alioth_static + kthread::kthread_static + krpc::rpc_static ${LLAMA_LIB} ${GGML_LIB} protobuf::libprotobuf diff --git a/kllm/config/arg.cc b/kllm/config/arg.cc index fbdea90..bde9a7d 100644 --- a/kllm/config/arg.cc +++ b/kllm/config/arg.cc @@ -32,7 +32,7 @@ #include #include #include -#include +#include #include namespace kllm { diff --git a/kllm/core/params.cc b/kllm/core/params.cc index 8f45d6c..d17d3bf 100644 --- a/kllm/core/params.cc +++ b/kllm/core/params.cc @@ -19,7 +19,7 @@ #include #include #include -#include +#include namespace kllm { diff --git a/kllm/openai/oai_processor.cc b/kllm/openai/oai_processor.cc index d85fc79..f91b7c8 100644 --- a/kllm/openai/oai_processor.cc +++ b/kllm/openai/oai_processor.cc @@ -646,46 +646,43 @@ namespace kllm { } - void setup_oai_api(KMContext *c) { + void setup_oai_api(KMContext *c, krpc::RestfulService *ins) { // APIS // register API routes - auto ins = krpc::RestfulService::instance(); - ins->set_processor("/health", std::shared_ptr(new HealthProcessor())); // public endpoint (no API key check) done - ins->set_processor("/metrics", std::shared_ptr(new MetricProcessor(c))); // done - ins->set_processor("/props", std::shared_ptr(new PropsProcessor(c))); // done - ins->set_processor("/models", std::shared_ptr(new ModelsProcessor(c))); // public endpoint (no API key check) done - ins->set_processor("/v1/models", std::shared_ptr(new ModelsProcessor(c))); // public endpoint (no API key check) done - ins->set_processor("/completion", std::shared_ptr(new CompletionsProcessor(c))); // legacy - ins->set_processor("/completions", std::shared_ptr(new CompletionsProcessor(c))); // done - ins->set_processor("/v1/completions", std::shared_ptr(new CompletionsProcessor(c))); // done - ins->set_processor("/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c))); // done - ins->set_processor("/v1/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c))); // done - ins->set_processor("/infill", std::shared_ptr(new InfillProcessor(c))); // done - ins->set_processor("/embedding", std::shared_ptr(new EmbeddingProcessor(c))); // legacy done - ins->set_processor("/embeddings", std::shared_ptr(new EmbeddingProcessor(c))); // done - ins->set_processor("/v1/embeddings", std::shared_ptr(new EmbeddingProcessor(c))); // done - ins->set_processor("/rerank", std::shared_ptr(new RerankProcessor(c))); // done - ins->set_processor("/reranking", std::shared_ptr(new RerankProcessor(c))); // done - ins->set_processor("/v1/rerank", std::shared_ptr(new RerankProcessor(c))); // done - ins->set_processor("/v1/reranking", std::shared_ptr(new RerankProcessor(c))); // done - ins->set_processor("/tokenize", std::shared_ptr(new TokenizeProcessor(c))); // done - ins->set_processor("/detokenize", std::shared_ptr(new DetokenizeProcessor(c))); // done + ins->set_processor("/health", std::shared_ptr(new HealthProcessor()), false); // public endpoint (no API key check) done + ins->set_processor("/metrics", std::shared_ptr(new MetricProcessor(c)), false); // done + ins->set_processor("/props", std::shared_ptr(new PropsProcessor(c)), false); // done + ins->set_processor("/models", std::shared_ptr(new ModelsProcessor(c)), false); // public endpoint (no API key check) done + ins->set_processor("/v1/models", std::shared_ptr(new ModelsProcessor(c)), false); // public endpoint (no API key check) done + ins->set_processor("/completion", std::shared_ptr(new CompletionsProcessor(c)), false); // legacy + ins->set_processor("/completions", std::shared_ptr(new CompletionsProcessor(c)), false); // done + ins->set_processor("/v1/completions", std::shared_ptr(new CompletionsProcessor(c)), false); // done + ins->set_processor("/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c)), false); // done + ins->set_processor("/v1/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c)), false); // done + ins->set_processor("/infill", std::shared_ptr(new InfillProcessor(c)), false); // done + ins->set_processor("/embedding", std::shared_ptr(new EmbeddingProcessor(c)), false); // legacy done + ins->set_processor("/embeddings", std::shared_ptr(new EmbeddingProcessor(c)), false); // done + ins->set_processor("/v1/embeddings", std::shared_ptr(new EmbeddingProcessor(c)), false); // done + ins->set_processor("/rerank", std::shared_ptr(new RerankProcessor(c)), false); // done + ins->set_processor("/reranking", std::shared_ptr(new RerankProcessor(c)), false); // done + ins->set_processor("/v1/rerank", std::shared_ptr(new RerankProcessor(c)), false); // done + ins->set_processor("/v1/reranking", std::shared_ptr(new RerankProcessor(c)), false); // done + ins->set_processor("/tokenize", std::shared_ptr(new TokenizeProcessor(c)), false); // done + ins->set_processor("/detokenize", std::shared_ptr(new DetokenizeProcessor(c)), false); // done // LoRA adapters hotswap - ins->set_processor("/lora-adapters", std::shared_ptr(new LoraListProcessor(c))); //done - ins->set_processor("/lora-adapters-apply", std::shared_ptr(new LoraApplyProcessor(c))); // done + ins->set_processor("/lora-adapters", std::shared_ptr(new LoraListProcessor(c)), false); //done + ins->set_processor("/lora-adapters-apply", std::shared_ptr(new LoraApplyProcessor(c)), false); // done // Save & load slots - ins->set_processor("/slots", std::shared_ptr(new SlotsProcessor(c))); // done + ins->set_processor("/slots", std::shared_ptr(new SlotsProcessor(c)), false); // done } - void setup_oai_ui(KMContext *c) { + void setup_oai_ui(KMContext *c, krpc::RestfulService *ins) { // register static assets routes - auto handle_static_content = [](const std::string &path, unsigned char *content, size_t len, const char *mime_type) { + auto handle_static_content = [ins](const std::string &path, unsigned char *content, size_t len, const char *mime_type) { std::string body(reinterpret_cast(content), len); - auto ins = krpc::RestfulService::instance(); - ins->set_static_content_processor(path,body, mime_type); + ins->set_static_content_processor(path,body, mime_type , false); }; auto public_path = c->params.public_path; - auto ins = krpc::RestfulService::instance(); if (!public_path.empty()) { // Set the base directory for serving static files auto sptr = new krpc::StaticFileProcessor(public_path); @@ -695,11 +692,10 @@ namespace kllm { if(public_path == "/") { ins->set_root_processor(std::shared_ptr(sptr)); } else { - ins->set_processor(public_path, std::shared_ptr(sptr)); + ins->set_processor(public_path, std::shared_ptr(sptr), false); } } else { // using embedded static files - auto ins = krpc::RestfulService::instance(); ins->set_root_processor(std::shared_ptr(new RootProcessor())); handle_static_content("/index.html", index_html, index_html_len, "text/html; charset=utf-8"); handle_static_content("/completion.js", completion_js, completion_js_len, "text/javascript; charset=utf-8"); diff --git a/kllm/openai/oai_processor.h b/kllm/openai/oai_processor.h index 2f60336..744e2b0 100644 --- a/kllm/openai/oai_processor.h +++ b/kllm/openai/oai_processor.h @@ -227,8 +227,8 @@ namespace kllm { void action(const krpc::RestfulRequest *request, krpc::RestfulResponse *response); }; - void setup_oai_api(KMContext *c); + void setup_oai_api(KMContext *c, krpc::RestfulService *ins); - void setup_oai_ui(KMContext *c); + void setup_oai_ui(KMContext *c, krpc::RestfulService *ins); } // namespace kllm diff --git a/kllm/tools/kllm.cc b/kllm/tools/kllm.cc index 79fae89..4677568 100644 --- a/kllm/tools/kllm.cc +++ b/kllm/tools/kllm.cc @@ -18,7 +18,7 @@ // A server to receive HttpRequest and send back HttpResponse. #include -#include +#include #include #include diff --git a/kllm/tools/oai/oai.cc b/kllm/tools/oai/oai.cc index 2b85616..744a220 100644 --- a/kllm/tools/oai/oai.cc +++ b/kllm/tools/oai/oai.cc @@ -19,7 +19,7 @@ #include #include #include -#include +#include #include #include #include @@ -51,11 +51,11 @@ namespace kllm { LOG_INF("%s: model loaded\n", __func__); krpc::Server server; - //auto *ins = krpc::RestfulService::instance(); + auto &ins = server.restful_service(); - setup_oai_api(&ServiceContext::ctx_server); + setup_oai_api(&ServiceContext::ctx_server, &ins); - setup_oai_ui(&ServiceContext::ctx_server); + setup_oai_ui(&ServiceContext::ctx_server, &ins); oai_context.start_context_async(); diff --git a/kmpkg-configuration.json b/kmpkg-configuration.json index 6d48a1c..0789c85 100644 --- a/kmpkg-configuration.json +++ b/kmpkg-configuration.json @@ -1,7 +1,7 @@ { "default-registry": { "kind": "git", - "baseline": "1c02bf9fc410b8fe6fe286c7b63083f05abe9087", + "baseline": "90bf767daaeaeba48d4c5e3cee23891f88c05d60", "repository": "https://gitee.com/kumo-pub/kmpkg" }, "registries": [ diff --git a/kmpkg.json b/kmpkg.json index 260f428..f915cf5 100644 --- a/kmpkg.json +++ b/kmpkg.json @@ -1,8 +1,8 @@ { "dependencies": [ - "cpp-httplib", "llama", "nlohmann-json", + "alioth", "krpc", { "name": "protobuf", -- Gitee From 6b8c2673b68012bd0f64b28140fe3345435f607d Mon Sep 17 00:00:00 2001 From: Jeff lothar Date: Sun, 19 Jan 2025 20:32:47 +0800 Subject: [PATCH 5/6] update krpc for auth control --- .gitignore | 3 +- cmake/kllm_deps.cmake | 1 - kllm/config/arg.cc | 2 +- kllm/core/params.cc | 2 +- kllm/openai/auth.cc | 60 +++++++++++++++++++++++++++ kllm/openai/auth.h | 43 +++++++++++++++++++ kllm/openai/oai_processor.cc | 78 ++++++++++++++++++++--------------- kllm/openai/oai_processor.h | 2 + kllm/public/completion.js | 2 + kllm/tools/service_context.cc | 1 + kmpkg-configuration.json | 2 +- kmpkg.json | 1 - 12 files changed, 157 insertions(+), 40 deletions(-) create mode 100644 kllm/openai/auth.cc create mode 100644 kllm/openai/auth.h diff --git a/.gitignore b/.gitignore index 521630c..d459d59 100644 --- a/.gitignore +++ b/.gitignore @@ -33,4 +33,5 @@ .idea cmake-build-debug build -.vscode \ No newline at end of file +.vscode +krpc_meta \ No newline at end of file diff --git a/cmake/kllm_deps.cmake b/cmake/kllm_deps.cmake index 804a401..fa0bc36 100644 --- a/cmake/kllm_deps.cmake +++ b/cmake/kllm_deps.cmake @@ -94,7 +94,6 @@ endif () set(KMCMAKE_DEPS_LINK turbo::turbo_static tally::tally_static - alioth::alioth_static kthread::kthread_static krpc::rpc_static ${LLAMA_LIB} diff --git a/kllm/config/arg.cc b/kllm/config/arg.cc index bde9a7d..f85ed0d 100644 --- a/kllm/config/arg.cc +++ b/kllm/config/arg.cc @@ -32,7 +32,7 @@ #include #include #include -#include +#include #include namespace kllm { diff --git a/kllm/core/params.cc b/kllm/core/params.cc index d17d3bf..208617e 100644 --- a/kllm/core/params.cc +++ b/kllm/core/params.cc @@ -19,7 +19,7 @@ #include #include #include -#include +#include namespace kllm { diff --git a/kllm/openai/auth.cc b/kllm/openai/auth.cc new file mode 100644 index 0000000..288e938 --- /dev/null +++ b/kllm/openai/auth.cc @@ -0,0 +1,60 @@ +// Copyright (C) 2024 Kumo inc. +// Author: Jeff.li lijippy@163.com +// All rights reserved. +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published +// by the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// + +#include +#include +namespace kllm { + + turbo::Status OaiAuth::init_oai_auth(krpc::AuthManager *manager) { + auto as = manager->get_auth_define(kGroup); + if(!as.ok()) { + krpc::AuthDefineBuild builder; + builder.set_group(kGroup) + .add_action(kCompletionAction) + .add_action(kEmbAction) + .add_collection(kCompletionCollection) + .add_collection(kEmbCollection); + krpc::AuthDefine def; + auto rs = builder.build(def); + if (!rs.ok()) { + return rs; + } + + rs = manager->create_auth_define(def); + if (!rs.ok()) { + return rs; + } + } + auto bs = manager->get_key("km_chat"); + if(!bs.ok()) { + krpc::ApiKeyBuilder b; + b.add_action(kGroup, kCompletionAction); + b.value("km_chat"); + b.add_collection(kGroup, kCompletionCollection); + b.expires(std::uint64_t(-1)); + b.auto_delete(false); + auto krs = b.build(manager); + if (!krs.ok()) { + return krs.status(); + } + return manager->create_key(krs.value()); + } + return turbo::OkStatus(); + } +} // namespace kllm + diff --git a/kllm/openai/auth.h b/kllm/openai/auth.h new file mode 100644 index 0000000..5b4d996 --- /dev/null +++ b/kllm/openai/auth.h @@ -0,0 +1,43 @@ +// Copyright (C) 2024 Kumo inc. +// Author: Jeff.li lijippy@163.com +// All rights reserved. +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published +// by the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . +// +// + +#pragma once + +#include +#include + +namespace kllm { + + class OaiAuth { + public: + static constexpr std::string_view kGroup = "oai"; + + static constexpr std::string_view kCompletionAction = "completion"; + static constexpr std::string_view kOaiCompletionAction = "oai:completion"; + static constexpr std::string_view kEmbAction = "emb"; + static constexpr std::string_view kAoiEmbAction = "oai:emb"; + + static constexpr std::string_view kCompletionCollection = "completion"; + static constexpr std::string_view kOaiCompletionCollection = "oai:completion"; + static constexpr std::string_view kEmbCollection = "emb"; + static constexpr std::string_view kAoiEmbCollection = "oai:emb"; + + static turbo::Status init_oai_auth(krpc::AuthManager *manager); + }; + +} // \ No newline at end of file diff --git a/kllm/openai/oai_processor.cc b/kllm/openai/oai_processor.cc index f91b7c8..8f80030 100644 --- a/kllm/openai/oai_processor.cc +++ b/kllm/openai/oai_processor.cc @@ -19,17 +19,19 @@ #include #include #include +#include #include -#include "kllm/public/index.html.h" -#include "kllm/public/completion.js.h" -#include "kllm/public/loading.html.h" -#include "kllm/public/deps_daisyui.min.css.h" -#include "kllm/public/deps_markdown-it.js.h" -#include "kllm/public/deps_tailwindcss.js.h" -#include "kllm/public/deps_vue.esm-browser.js.h" -#include "kllm/public/favicon.png.h" -#include "kllm/public/kumo_logo.svg.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include #define MIMETYPE_JSON "application/json; charset=utf-8" @@ -649,38 +651,46 @@ namespace kllm { void setup_oai_api(KMContext *c, krpc::RestfulService *ins) { // APIS // register API routes - ins->set_processor("/health", std::shared_ptr(new HealthProcessor()), false); // public endpoint (no API key check) done - ins->set_processor("/metrics", std::shared_ptr(new MetricProcessor(c)), false); // done - ins->set_processor("/props", std::shared_ptr(new PropsProcessor(c)), false); // done - ins->set_processor("/models", std::shared_ptr(new ModelsProcessor(c)), false); // public endpoint (no API key check) done - ins->set_processor("/v1/models", std::shared_ptr(new ModelsProcessor(c)), false); // public endpoint (no API key check) done - ins->set_processor("/completion", std::shared_ptr(new CompletionsProcessor(c)), false); // legacy - ins->set_processor("/completions", std::shared_ptr(new CompletionsProcessor(c)), false); // done - ins->set_processor("/v1/completions", std::shared_ptr(new CompletionsProcessor(c)), false); // done - ins->set_processor("/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c)), false); // done - ins->set_processor("/v1/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c)), false); // done - ins->set_processor("/infill", std::shared_ptr(new InfillProcessor(c)), false); // done - ins->set_processor("/embedding", std::shared_ptr(new EmbeddingProcessor(c)), false); // legacy done - ins->set_processor("/embeddings", std::shared_ptr(new EmbeddingProcessor(c)), false); // done - ins->set_processor("/v1/embeddings", std::shared_ptr(new EmbeddingProcessor(c)), false); // done - ins->set_processor("/rerank", std::shared_ptr(new RerankProcessor(c)), false); // done - ins->set_processor("/reranking", std::shared_ptr(new RerankProcessor(c)), false); // done - ins->set_processor("/v1/rerank", std::shared_ptr(new RerankProcessor(c)), false); // done - ins->set_processor("/v1/reranking", std::shared_ptr(new RerankProcessor(c)), false); // done - ins->set_processor("/tokenize", std::shared_ptr(new TokenizeProcessor(c)), false); // done - ins->set_processor("/detokenize", std::shared_ptr(new DetokenizeProcessor(c)), false); // done + auto m = ins->get_auth(); + if(!m) { + m= ins->create_default_auth("kllm_boost_key"); + } + auto rs = OaiAuth::init_oai_auth(m.get()); + if(!rs.ok()) { + LOG(FATAL)<<"init_oai_auth fail"<set_processor("GET", "/health", std::shared_ptr(new HealthProcessor()), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // public endpoint (no API key check) done + ins->set_processor("GET", "/metrics", std::shared_ptr(new MetricProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("GET", "/props", std::shared_ptr(new PropsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("GET", "/models", std::shared_ptr(new ModelsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // public endpoint (no API key check) done + ins->set_processor("GET","/v1/models", std::shared_ptr(new ModelsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // public endpoint (no API key check) done + ins->set_processor("POST", "/completion", std::shared_ptr(new CompletionsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // legacy + ins->set_processor("POST","/completions", std::shared_ptr(new CompletionsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/v1/completions", std::shared_ptr(new CompletionsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c)), OaiAuth::kOaiCompletionAction, OaiAuth::kOaiCompletionCollection); // done + ins->set_processor("POST","/v1/chat/completions", std::shared_ptr(new ChatCompletionsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/infill", std::shared_ptr(new InfillProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/embedding", std::shared_ptr(new EmbeddingProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // legacy done + ins->set_processor("POST","/embeddings", std::shared_ptr(new EmbeddingProcessor(c)),krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/v1/embeddings", std::shared_ptr(new EmbeddingProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/rerank", std::shared_ptr(new RerankProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/reranking", std::shared_ptr(new RerankProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/v1/rerank", std::shared_ptr(new RerankProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/v1/reranking", std::shared_ptr(new RerankProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/tokenize", std::shared_ptr(new TokenizeProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done + ins->set_processor("POST","/detokenize", std::shared_ptr(new DetokenizeProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done // LoRA adapters hotswap - ins->set_processor("/lora-adapters", std::shared_ptr(new LoraListProcessor(c)), false); //done - ins->set_processor("/lora-adapters-apply", std::shared_ptr(new LoraApplyProcessor(c)), false); // done + ins->set_processor("GET", "/lora-adapters", std::shared_ptr(new LoraListProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); //done + ins->set_processor("POST", "/lora-adapters-apply", std::shared_ptr(new LoraApplyProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done // Save & load slots - ins->set_processor("/slots", std::shared_ptr(new SlotsProcessor(c)), false); // done + ins->set_processor("GET","/slots", std::shared_ptr(new SlotsProcessor(c)), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); // done } void setup_oai_ui(KMContext *c, krpc::RestfulService *ins) { // register static assets routes auto handle_static_content = [ins](const std::string &path, unsigned char *content, size_t len, const char *mime_type) { std::string body(reinterpret_cast(content), len); - ins->set_static_content_processor(path,body, mime_type , false); + ins->set_static_content_processor(path,body, mime_type); }; auto public_path = c->params.public_path; if (!public_path.empty()) { @@ -692,7 +702,7 @@ namespace kllm { if(public_path == "/") { ins->set_root_processor(std::shared_ptr(sptr)); } else { - ins->set_processor(public_path, std::shared_ptr(sptr), false); + ins->set_processor("GET", public_path, std::shared_ptr(sptr), krpc::ApiKeyOperation::kPublicAction, krpc::ApiKeyOperation::kPublicCollections); } } else { // using embedded static files diff --git a/kllm/openai/oai_processor.h b/kllm/openai/oai_processor.h index 744e2b0..e2b5423 100644 --- a/kllm/openai/oai_processor.h +++ b/kllm/openai/oai_processor.h @@ -18,6 +18,7 @@ #pragma once #include +#include #include namespace kllm { @@ -29,6 +30,7 @@ namespace kllm { void process(const krpc::RestfulRequest *request, krpc::RestfulResponse *response) override; bool is_wildcards() const override; + }; struct MetricProcessor : public krpc::RestfulProcessor { diff --git a/kllm/public/completion.js b/kllm/public/completion.js index 54a0f22..80e327f 100644 --- a/kllm/public/completion.js +++ b/kllm/public/completion.js @@ -40,6 +40,8 @@ export async function* llama(prompt, params = {}, config = {}) { 'Connection': 'keep-alive', 'Content-Type': 'application/json', 'Accept': 'text/event-stream', + 'x-kumo-key':'km_chat', + //...(params.api_key ? {'x-kumo-key': `${params.api_key}`} : {}) ...(params.api_key ? {'Authorization': `Bearer ${params.api_key}`} : {}) }, signal: controller.signal, diff --git a/kllm/tools/service_context.cc b/kllm/tools/service_context.cc index eeac944..d42b973 100644 --- a/kllm/tools/service_context.cc +++ b/kllm/tools/service_context.cc @@ -8,6 +8,7 @@ #include #include #include +#include #include namespace kllm { diff --git a/kmpkg-configuration.json b/kmpkg-configuration.json index 0789c85..469d048 100644 --- a/kmpkg-configuration.json +++ b/kmpkg-configuration.json @@ -1,7 +1,7 @@ { "default-registry": { "kind": "git", - "baseline": "90bf767daaeaeba48d4c5e3cee23891f88c05d60", + "baseline": "93fd707693c7337790ad1f61fe8bc751a4a5b239", "repository": "https://gitee.com/kumo-pub/kmpkg" }, "registries": [ diff --git a/kmpkg.json b/kmpkg.json index f915cf5..9a68470 100644 --- a/kmpkg.json +++ b/kmpkg.json @@ -2,7 +2,6 @@ "dependencies": [ "llama", "nlohmann-json", - "alioth", "krpc", { "name": "protobuf", -- Gitee From 0b780b66f28402621c042c41af149fa44aceec7c Mon Sep 17 00:00:00 2001 From: Jeff lothar Date: Sun, 19 Jan 2025 20:58:00 +0800 Subject: [PATCH 6/6] fix auth permission --- kllm/openai/oai_processor.h | 62 +++++++++++++++++++++++++++---------- 1 file changed, 46 insertions(+), 16 deletions(-) diff --git a/kllm/openai/oai_processor.h b/kllm/openai/oai_processor.h index e2b5423..5f6f0b2 100644 --- a/kllm/openai/oai_processor.h +++ b/kllm/openai/oai_processor.h @@ -20,11 +20,39 @@ #include #include #include +#include +#include namespace kllm { - struct HealthProcessor : public krpc::RestfulProcessor { + struct OaiProcessorBase : public krpc::RestfulProcessor { + turbo::Status get_auth_authed(const krpc::RestfulRequest *request, std::string &result) final { + static std::string_view oai_prefix = "Bearer"; + auto a = request->get_authorization(); + auto api_key = request->get_api_key(); + std::string_view av; + if (a) { + av = *a; + } else if (!api_key.empty()) { + av = api_key; + } + if (!av.empty()) { + av = turbo::trim_all(av); + if (turbo::starts_with_ignore_case(av, oai_prefix)) { + av.remove_prefix(oai_prefix.size()); + av = turbo::trim_all(av); + } + result = av; + LOG(INFO)<