Ling-3.0-tiny-RKNN / include /ling3 /chat_protocol.h
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
12.7 kB
#pragma once
#include <nlohmann/json.hpp>
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <limits>
#include <stdexcept>
#include <string>
#include <vector>
namespace ling3::chat {
using Json = nlohmann::json;
inline constexpr const char * kModel = "mindnano-ling3-tiny";
struct Error : std::runtime_error {
int status;
std::string code;
Error(int s, std::string c, std::string m)
: std::runtime_error(std::move(m)), status(s), code(std::move(c)) {}
};
inline Json ErrorBody(const Error & e) {
return {{"error", {{"message", e.what()}, {"type", "invalid_request_error"},
{"param", nullptr}, {"code", e.code}}}};
}
struct Request {
std::string prompt;
std::size_t max_tokens = std::numeric_limits<std::size_t>::max();
bool stream = false, include_usage = false;
bool cache_prompt = true;
bool reuse_generated_state = true;
std::string cache_user;
std::string session_id;
double temperature = 0, top_p = 1;
int top_k = 0;
double repeat_penalty = 1;
bool enable_thinking = false, flow_ack = false;
std::uint32_t seed = 0;
std::vector<std::string> stops;
};
inline std::size_t OutputBudget(std::size_t prompt, std::size_t capacity,
std::size_t requested) {
if (!prompt || prompt > capacity)
throw Error(400, "context_length_exceeded", "input exceeds allocated context of " +
std::to_string(capacity) + " tokens, or is empty");
return std::min(requested, capacity - prompt);
}
inline Request Parse(const Json & j) {
auto fail = [](std::string m) { throw Error(400, "invalid_parameter", std::move(m)); };
if (!j.is_object()) fail("request must be a JSON object");
for (auto it = j.begin(); it != j.end(); ++it) {
const std::vector<std::string> supported {"model", "messages", "max_tokens",
"max_completion_tokens", "stream", "stream_options", "temperature", "top_p",
"seed", "stop", "n", "user", "cache_prompt", "top_k", "repeat_penalty",
"enable_thinking", "flow_control", "session_id", "reuse_generated_state"};
if (std::find(supported.begin(), supported.end(), it.key()) == supported.end())
fail("unsupported parameter: " + it.key());
}
if (!j.contains("model") || !j["model"].is_string() || j["model"] != kModel)
throw Error(404, "model_not_found", std::string("model must be ") + kModel);
Request r;
if (j.contains("reuse_generated_state")) {
if (!j["reuse_generated_state"].is_boolean()) fail("reuse_generated_state must be boolean");
r.reuse_generated_state = j["reuse_generated_state"].get<bool>();
}
if (j.contains("session_id")) {
if (!j["session_id"].is_string()) fail("session_id must be a string");
r.session_id = j["session_id"].get<std::string>();
if (r.session_id.empty() || r.session_id.size() > 256) fail("session_id must contain 1..256 bytes");
}
if (j.contains("top_k")) {
if (!j["top_k"].is_number_integer() || j["top_k"] < 0 || j["top_k"] > 1000000)
fail("top_k must be an integer in 0..1000000 (0 disables it)");
r.top_k = j["top_k"].get<int>();
}
if (j.contains("repeat_penalty")) {
if (!j["repeat_penalty"].is_number()) fail("repeat_penalty must be numeric");
r.repeat_penalty = j["repeat_penalty"].get<double>();
if (!std::isfinite(r.repeat_penalty) || r.repeat_penalty <= 0)
fail("repeat_penalty must be finite and positive");
}
if (j.contains("enable_thinking")) {
if (!j["enable_thinking"].is_boolean()) fail("enable_thinking must be boolean");
r.enable_thinking = j["enable_thinking"].get<bool>();
}
if (j.contains("flow_control")) {
if (j["flow_control"] != "none" && j["flow_control"] != "ack") fail("invalid flow_control");
r.flow_ack = j["flow_control"] == "ack";
if (r.flow_ack && j.value("stream", false) != true) fail("flow_control=ack requires streaming");
}
if (j.contains("cache_prompt")) {
if (!j["cache_prompt"].is_boolean()) fail("cache_prompt must be boolean");
r.cache_prompt = j["cache_prompt"].get<bool>();
}
if (j.contains("max_tokens") && j.contains("max_completion_tokens"))
fail("set only one of max_tokens and max_completion_tokens");
for (auto key : {"max_tokens", "max_completion_tokens"}) {
if (!j.contains(key)) continue;
const auto & value = j[key];
if (!value.is_number_integer() ||
(value.is_number_unsigned() ? value.get<std::uint64_t>() == 0
: value.get<std::int64_t>() <= 0))
fail("max_tokens/max_completion_tokens must be a positive integer");
const auto count = value.get<std::uint64_t>();
if (count > std::numeric_limits<std::size_t>::max()) fail("output budget exceeds addressable integer range");
r.max_tokens = static_cast<std::size_t>(count);
}
if (j.contains("stream")) {
if (!j["stream"].is_boolean()) fail("stream must be boolean");
r.stream = j["stream"].get<bool>();
}
if (j.contains("stream_options") && !j["stream_options"].is_null()) {
const auto & s = j["stream_options"];
if (!r.stream || !s.is_object()) fail("stream_options requires stream=true");
for (auto it = s.begin(); it != s.end(); ++it)
if (it.key() != "include_usage" || !it.value().is_boolean())
fail("only boolean stream_options.include_usage is supported");
r.include_usage = s.value("include_usage", false);
}
for (auto key : {"temperature", "top_p"}) {
if (!j.contains(key)) continue;
if (!j[key].is_number()) fail(std::string(key) + " must be numeric");
const auto x = j[key].get<double>();
if (!std::isfinite(x) || x < 0 || x > (std::string(key) == "temperature" ? 2 : 1))
fail(std::string(key) + " out of range");
if (std::string(key) == "temperature") r.temperature = x;
else { if (x == 0) fail("top_p must be > 0"); r.top_p = x; }
}
if (j.contains("seed")) {
if (!j["seed"].is_number_integer() || j["seed"] < 0 || j["seed"] > UINT32_MAX)
fail("seed must be a uint32 integer");
r.seed = j["seed"].get<std::uint32_t>();
}
if (j.contains("n") && (!j["n"].is_number_integer() || j["n"] != 1)) fail("only n=1 is supported");
if (j.contains("user") && !j["user"].is_string()) fail("user must be a string");
if (j.contains("user")) r.cache_user = j["user"].get<std::string>();
if (j.contains("stop") && !j["stop"].is_null()) {
const auto s = j["stop"].is_string() ? Json::array({j["stop"]}) : j["stop"];
if (!s.is_array() || s.size() > 4) fail("stop must be a string or at most 4 strings");
for (const auto & v : s) {
if (!v.is_string() || v.get_ref<const std::string &>().empty() ||
v.get_ref<const std::string &>().size() > 256) fail("invalid stop string");
r.stops.push_back(v.get<std::string>());
}
}
if (!j.contains("messages") || !j["messages"].is_array() || j["messages"].empty())
fail("messages must be a non-empty array");
const auto & messages = j["messages"];
for (const auto & m : messages) {
if (!m.is_object() || !m.contains("role") || !m["role"].is_string() ||
!m.contains("content") || !m["content"].is_string())
fail("messages require string role and content; text only");
const auto role = m["role"].get<std::string>();
if (role != "system" && role != "user" && role != "assistant") fail("unsupported message role");
for (auto it = m.begin(); it != m.end(); ++it)
if (it.key() != "role" && it.key() != "content") fail("unsupported message field: " + it.key());
}
if (messages.back()["role"] != "user") fail("last message must be a user message");
// Bailing V3 text template. No hidden, server-owned conversation.
const std::string thinking = r.enable_thinking ? "on" : "off";
r.prompt = "<role>SYSTEM</role>";
std::size_t start = 0;
if (messages.front()["role"] == "system") {
auto s = messages.front()["content"].get<std::string>();
// The explicit API switch owns template control. Normalize conflicting
// legacy instructions rather than producing two contradictory modes.
for (const auto old : {std::string("detailed thinking off"), std::string("detailed thinking on")}) {
std::size_t at = 0;
while ((at = s.find(old, at)) != std::string::npos) {
const auto replacement = "detailed thinking " + thinking;
s.replace(at, old.size(), replacement); at += replacement.size();
}
}
r.prompt += s;
if (s.find("detailed thinking off") == std::string::npos && s.find("detailed thinking on") == std::string::npos)
r.prompt += "\ndetailed thinking " + thinking;
start = 1;
} else r.prompt += "detailed thinking " + thinking;
r.prompt += "<|role_end|>";
for (std::size_t i = start; i < messages.size(); ++i) {
const auto role = messages[i]["role"].get<std::string>();
auto content = messages[i]["content"].get<std::string>();
if (role == "assistant") {
const auto close = content.find("</think>");
if (close == std::string::npos) r.prompt += "<role>ASSISTANT</role>\n<think></think>";
else {
auto reasoning = content.substr(0, close);
const auto open = reasoning.rfind("<think>");
if (open != std::string::npos) reasoning.erase(0, open + 7);
const auto first = reasoning.find_first_not_of('\n');
reasoning = first == std::string::npos ? "" : reasoning.substr(first, reasoning.find_last_not_of('\n') - first + 1);
content = content.substr(content.rfind("</think>") + 8);
content.erase(0, std::min(content.find_first_not_of('\n'), content.size()));
r.prompt += "<role>ASSISTANT</role>\n<think>" + reasoning + "</think>";
}
}
else r.prompt += role == "user" ? "<role>HUMAN</role>" : "<role>SYSTEM</role>";
r.prompt += content + "<|role_end|>";
}
r.prompt += r.enable_thinking ? "<role>ASSISTANT</role>\n<think>" : "<role>ASSISTANT</role>\n<think></think>";
return r;
}
// Token pieces can split a UTF-8 character. Keep incomplete sequences until the
// next token; JSON encoders must never see partial bytes.
inline std::string TakeUtf8(std::string & pending, bool flush) {
std::string out;
std::size_t i = 0;
while (i < pending.size()) {
const auto c = static_cast<unsigned char>(pending[i]);
std::size_t n = c < 0x80 ? 1 : c >= 0xc2 && c <= 0xdf ? 2 :
c >= 0xe0 && c <= 0xef ? 3 : c >= 0xf0 && c <= 0xf4 ? 4 : 0;
if (n && i + n > pending.size() && !flush) break;
bool valid = n && i + n <= pending.size();
for (std::size_t k = 1; valid && k < n; ++k)
valid = (static_cast<unsigned char>(pending[i + k]) & 0xc0) == 0x80;
if (valid && n >= 3) {
const auto b = static_cast<unsigned char>(pending[i + 1]);
valid = !((c == 0xe0 && b < 0xa0) || (c == 0xed && b >= 0xa0) ||
(c == 0xf0 && b < 0x90) || (c == 0xf4 && b >= 0x90));
}
if (valid) { out.append(pending, i, n); i += n; }
else { out += "\xef\xbf\xbd"; ++i; }
}
pending.erase(0, i);
return out;
}
class TextFilter {
std::string pending_, utf8_;
std::vector<std::string> stops_;
public:
bool stopped = false;
explicit TextFilter(std::vector<std::string> stops) : stops_(std::move(stops)) {}
std::string Push(std::string_view piece, bool final = false) {
if (stopped) return {};
pending_ += piece;
std::size_t end = std::string::npos;
for (const auto & stop : stops_) end = std::min(end, pending_.find(stop));
if (end != std::string::npos) { pending_.resize(end); stopped = true; final = true; }
std::size_t keep = 0;
if (!final) for (const auto & stop : stops_)
for (std::size_t k = 1; k < stop.size() && k <= pending_.size(); ++k)
if (pending_.compare(pending_.size() - k, k, stop, 0, k) == 0) keep = std::max(keep, k);
utf8_ += pending_.substr(0, pending_.size() - keep);
pending_.erase(0, pending_.size() - keep);
return TakeUtf8(utf8_, final);
}
};
} // namespace ling3::chat