#include "ling3/service.h" #include "ling3/decoder.h" #include "ling3/model_package.h" #include "ling3/tokenizer.h" #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #include #if defined(__linux__) #include #include #include #include #include #include #endif namespace ling3 { namespace { using Clock = std::chrono::steady_clock; double Milliseconds(Clock::time_point begin, Clock::time_point end) { return std::chrono::duration(end - begin).count(); } template std::span TensorSpan(const TensorView & tensor) { if (tensor.entry->data_bytes % sizeof(T) != 0) { throw std::runtime_error(std::string(tensor.name) + " has an incompatible byte count"); } return { reinterpret_cast(tensor.data), static_cast(tensor.entry->data_bytes / sizeof(T)), }; } Tokenizer PackageTokenizer(const ModelPackage & package) { const auto & asset = package.tensor("tokenizer"); if (asset.entry->role != static_cast(TensorRole::kTokenizer)) { throw std::runtime_error("tokenizer tensor has the wrong role"); } return Tokenizer(TensorSpan(asset)); } std::string JsonEscape(std::string_view input) { std::string output; output.reserve(input.size() + 8); constexpr char hex[] = "0123456789abcdef"; for (const unsigned char value : input) { switch (value) { case '"': output += "\\\""; break; case '\\': output += "\\\\"; break; case '\b': output += "\\b"; break; case '\f': output += "\\f"; break; case '\n': output += "\\n"; break; case '\r': output += "\\r"; break; case '\t': output += "\\t"; break; default: if (value < 0x20) { output += "\\u00"; output.push_back(hex[value >> 4U]); output.push_back(hex[value & 0x0FU]); } else { output.push_back(static_cast(value)); } } } return output; } std::string JsonNumber(double value) { std::ostringstream stream; stream << std::fixed << std::setprecision(3) << value; return stream.str(); } class SessionEngine { public: SessionEngine( const ModelPackage & package, const Tokenizer & tokenizer, Decoder & decoder, double initialization_ms) : package_(package), tokenizer_(tokenizer), decoder_(decoder), logits_(package.header().vocab_size), initialization_ms_(initialization_ms) {} bool active() const noexcept { return active_.load(std::memory_order_relaxed); } std::size_t position() const noexcept { return position_.load(std::memory_order_relaxed); } std::size_t turns() const noexcept { return turns_.load(std::memory_order_relaxed); } double initialization_ms() const noexcept { return initialization_ms_; } bool Cancel() noexcept { if (!active()) return false; cancel_requested_.store(true, std::memory_order_relaxed); return true; } void Reset() { cancel_requested_.store(true, std::memory_order_relaxed); std::lock_guard lock(inference_mutex_); decoder_.Reset(); position_.store(0, std::memory_order_relaxed); turns_.store(0, std::memory_order_relaxed); cancel_requested_.store(false, std::memory_order_relaxed); } void Chat( std::string_view user_text, std::size_t maximum_new_tokens, const std::function & emit) { const auto request_begin = Clock::now(); std::unique_lock lock(inference_mutex_); const auto lock_acquired = Clock::now(); cancel_requested_.store(false, std::memory_order_relaxed); active_.store(true, std::memory_order_relaxed); struct ActiveGuard { std::atomic_bool & active; ~ActiveGuard() { active.store(false, std::memory_order_relaxed); } } guard {active_}; if (user_text.empty()) throw std::invalid_argument("request body must contain user text"); if (maximum_new_tokens == 0 || maximum_new_tokens > 2048) { throw std::invalid_argument("max_tokens must be in [1, 2048]"); } const bool first_turn = decoder_.position() == 0; const std::string prompt = first_turn ? "SYSTEMdetailed thinking off<|role_end|>" "HUMAN" + std::string(user_text) + "<|role_end|>ASSISTANT\n" : "<|role_end|>HUMAN" + std::string(user_text) + "<|role_end|>ASSISTANT\n"; const auto prompt_tokens = tokenizer_.Encode(prompt); if (prompt_tokens.empty() || decoder_.position() + prompt_tokens.size() + maximum_new_tokens > package_.header().max_context) { throw std::runtime_error("request would exceed the package context capacity"); } const std::string start_event = "{\"type\":\"start\",\"turn\":" + std::to_string(turns() + 1) + ",\"position\":" + std::to_string(decoder_.position()) + ",\"prompt_tokens\":" + std::to_string(prompt_tokens.size()) + ",\"queue_ms\":" + JsonNumber(Milliseconds(request_begin, lock_acquired)) + "}\n"; if (!emit(start_event)) return; const auto prefill_begin = Clock::now(); double prefill_eval_ms = 0.0; std::size_t prefill_offset = 0; std::size_t prefill_batch_count = 0; std::size_t batch32_count = 0; std::size_t batched_tokens = 0; std::size_t maximum_batch_rows = 0; const std::size_t batch_granularity = decoder_.batch_granularity(); while (batch_granularity != 0 && prompt_tokens.size() - prefill_offset >= 2) { std::size_t rows = std::min( prompt_tokens.size() - prefill_offset, 128); rows -= rows % batch_granularity; if (rows < 2) break; const auto block = std::span(prompt_tokens).subspan( prefill_offset, rows); const bool needs_logits = prompt_tokens.size() - prefill_offset == rows; prefill_eval_ms += (needs_logits ? decoder_.EvalBatch(block, logits_) : decoder_.EvalBatchState(block)).total_ms; prefill_offset += rows; ++prefill_batch_count; batch32_count += rows == 32 ? 1 : 0; batched_tokens += rows; maximum_batch_rows = std::max(maximum_batch_rows, rows); position_.store(decoder_.position(), std::memory_order_relaxed); } for (; prefill_offset < prompt_tokens.size(); ++prefill_offset) { prefill_eval_ms += decoder_.Eval(prompt_tokens[prefill_offset], logits_).total_ms; position_.store(decoder_.position(), std::memory_order_relaxed); } const auto prefill_end = Clock::now(); std::size_t generated_tokens = 0; double decode_eval_ms = 0.0; bool stopped_on_eos = false; bool client_connected = true; const auto generation_begin = Clock::now(); double first_token_ms = 0.0; for (std::size_t index = 0; index < maximum_new_tokens; ++index) { if (cancel_requested_.load(std::memory_order_relaxed)) break; const auto found = std::max_element(logits_.begin(), logits_.end()); const auto token = static_cast(found - logits_.begin()); if (token == package_.header().eos_token) { stopped_on_eos = true; break; } if (generated_tokens == 0) { first_token_ms = Milliseconds(request_begin, Clock::now()); } const std::string event = "{\"type\":\"token\",\"id\":" + std::to_string(token) + ",\"text\":\"" + JsonEscape(tokenizer_.Piece(token)) + "\",\"index\":" + std::to_string(generated_tokens) + ",\"elapsed_ms\":" + JsonNumber(Milliseconds(request_begin, Clock::now())) + "}\n"; if (!emit(event)) { client_connected = false; cancel_requested_.store(true, std::memory_order_relaxed); break; } ++generated_tokens; const auto timing = decoder_.Eval(token, logits_); decode_eval_ms += timing.total_ms; position_.store(decoder_.position(), std::memory_order_relaxed); } turns_.fetch_add(1, std::memory_order_relaxed); const bool canceled = cancel_requested_.load(std::memory_order_relaxed); const double prefill_ms = Milliseconds(prefill_begin, prefill_end); const double generation_ms = Milliseconds(generation_begin, Clock::now()); if (client_connected) { const std::string finish_event = "{\"type\":\"finish\",\"reason\":\"" + std::string(canceled ? "canceled" : (stopped_on_eos ? "eos" : "length")) + "\",\"turn\":" + std::to_string(turns()) + ",\"position\":" + std::to_string(decoder_.position()) + ",\"prompt_tokens\":" + std::to_string(prompt_tokens.size()) + ",\"prefill_batch_count\":" + std::to_string(prefill_batch_count) + ",\"prefill_batch32_count\":" + std::to_string(batch32_count) + ",\"prefill_batched_tokens\":" + std::to_string(batched_tokens) + ",\"prefill_max_batch_rows\":" + std::to_string(maximum_batch_rows) + ",\"generated_tokens\":" + std::to_string(generated_tokens) + ",\"queue_ms\":" + JsonNumber(Milliseconds(request_begin, lock_acquired)) + ",\"prefill_ms\":" + JsonNumber(prefill_ms) + ",\"prefill_eval_ms\":" + JsonNumber(prefill_eval_ms) + ",\"prefill_tokens_per_second\":" + JsonNumber(prefill_ms == 0.0 ? 0.0 : 1000.0 * prompt_tokens.size() / prefill_ms) + ",\"ttft_ms\":" + JsonNumber(first_token_ms) + ",\"generation_ms\":" + JsonNumber(generation_ms) + ",\"decode_eval_ms\":" + JsonNumber(decode_eval_ms) + ",\"decode_tokens_per_second\":" + JsonNumber(decode_eval_ms == 0.0 ? 0.0 : 1000.0 * generated_tokens / decode_eval_ms) + "}\n"; emit(finish_event); } } private: const ModelPackage & package_; const Tokenizer & tokenizer_; Decoder & decoder_; std::vector logits_; std::mutex inference_mutex_; std::atomic_bool active_ {false}; std::atomic_bool cancel_requested_ {false}; std::atomic_size_t position_ {0}; std::atomic_size_t turns_ {0}; double initialization_ms_ = 0.0; }; #if defined(__linux__) struct HttpRequest { std::string method; std::string target; std::string body; }; class Socket { public: explicit Socket(int descriptor = -1) : descriptor_(descriptor) {} ~Socket() { if (descriptor_ >= 0) close(descriptor_); } Socket(const Socket &) = delete; Socket & operator=(const Socket &) = delete; Socket(Socket && other) noexcept : descriptor_(std::exchange(other.descriptor_, -1)) {} int get() const noexcept { return descriptor_; } private: int descriptor_; }; bool SendAll(int socket, std::string_view data) { while (!data.empty()) { const auto sent = send(socket, data.data(), data.size(), MSG_NOSIGNAL); if (sent < 0 && errno == EINTR) continue; if (sent <= 0) return false; data.remove_prefix(static_cast(sent)); } return true; } void SendResponse( int socket, int status, std::string_view reason, std::string_view body, std::string_view content_type = "application/json; charset=utf-8") { const std::string header = "HTTP/1.1 " + std::to_string(status) + " " + std::string(reason) + "\r\n" "Content-Type: " + std::string(content_type) + "\r\n" "Content-Length: " + std::to_string(body.size()) + "\r\n" "Connection: close\r\n\r\n"; SendAll(socket, header); SendAll(socket, body); } bool SendChunk(int socket, std::string_view data) { std::ostringstream size; size << std::hex << data.size() << "\r\n"; return SendAll(socket, size.str()) && SendAll(socket, data) && SendAll(socket, "\r\n"); } std::size_t ParseContentLength(std::string_view headers) { std::size_t cursor = headers.find("\r\n") + 2; while (cursor < headers.size()) { const auto end = headers.find("\r\n", cursor); if (end == std::string_view::npos || end == cursor) break; const auto line = headers.substr(cursor, end - cursor); const auto colon = line.find(':'); if (colon != std::string_view::npos) { std::string name(line.substr(0, colon)); std::transform(name.begin(), name.end(), name.begin(), [](unsigned char value) { return static_cast(std::tolower(value)); }); if (name == "content-length") { auto value = line.substr(colon + 1); while (!value.empty() && value.front() == ' ') value.remove_prefix(1); std::size_t length = 0; const auto [end_ptr, error] = std::from_chars(value.data(), value.data() + value.size(), length); if (error != std::errc {} || end_ptr != value.data() + value.size()) { throw std::runtime_error("invalid Content-Length"); } return length; } } cursor = end + 2; } return 0; } HttpRequest ReadRequest(int socket) { constexpr std::size_t kMaximumHeaders = 64 * 1024; constexpr std::size_t kMaximumBody = 1024 * 1024; std::string data; data.reserve(4096); std::size_t header_end = std::string::npos; char buffer[4096]; while ((header_end = data.find("\r\n\r\n")) == std::string::npos) { const auto count = recv(socket, buffer, sizeof(buffer), 0); if (count < 0 && errno == EINTR) continue; if (count <= 0) throw std::runtime_error("connection closed before request headers"); data.append(buffer, static_cast(count)); if (data.size() > kMaximumHeaders) throw std::runtime_error("request headers are too large"); } header_end += 4; const std::string_view headers(data.data(), header_end); const auto request_line_end = headers.find("\r\n"); const auto first_space = headers.find(' '); const auto second_space = first_space == std::string_view::npos ? std::string_view::npos : headers.find(' ', first_space + 1); if (request_line_end == std::string_view::npos || first_space == std::string_view::npos || second_space == std::string_view::npos || second_space > request_line_end) { throw std::runtime_error("invalid HTTP request line"); } HttpRequest request; request.method = std::string(headers.substr(0, first_space)); request.target = std::string(headers.substr(first_space + 1, second_space - first_space - 1)); const auto body_bytes = ParseContentLength(headers); if (body_bytes > kMaximumBody) throw std::runtime_error("request body is too large"); while (data.size() - header_end < body_bytes) { const auto count = recv(socket, buffer, sizeof(buffer), 0); if (count < 0 && errno == EINTR) continue; if (count <= 0) throw std::runtime_error("connection closed before request body"); data.append(buffer, static_cast(count)); } request.body.assign(data.data() + header_end, body_bytes); return request; } std::string_view Path(std::string_view target) { const auto query = target.find('?'); return target.substr(0, query); } std::size_t MaxTokens(std::string_view target) { constexpr std::string_view key = "max_tokens="; const auto found = target.find(key); if (found == std::string_view::npos) return 64; const auto begin = found + key.size(); const auto end = target.find('&', begin); const auto value = target.substr(begin, end == std::string_view::npos ? end : end - begin); std::size_t result = 0; const auto [end_ptr, error] = std::from_chars(value.data(), value.data() + value.size(), result); if (error != std::errc {} || end_ptr != value.data() + value.size()) { throw std::invalid_argument("invalid max_tokens query parameter"); } return result; } void HandleClient(Socket client, SessionEngine & engine) { try { const auto request = ReadRequest(client.get()); const auto path = Path(request.target); if (request.method == "GET" && path == "/health") { const std::string body = "{\"status\":\"ok\",\"active\":" + std::string(engine.active() ? "true" : "false") + ",\"position\":" + std::to_string(engine.position()) + ",\"turns\":" + std::to_string(engine.turns()) + ",\"initialization_ms\":" + JsonNumber(engine.initialization_ms()) + "}\n"; SendResponse(client.get(), 200, "OK", body); return; } if (request.method == "POST" && path == "/v1/cancel") { const bool canceled = engine.Cancel(); SendResponse(client.get(), 200, "OK", canceled ? "{\"cancel_requested\":true}\n" : "{\"cancel_requested\":false}\n"); return; } if (request.method == "POST" && path == "/v1/reset") { engine.Reset(); SendResponse(client.get(), 200, "OK", "{\"reset\":true}\n"); return; } if (request.method == "POST" && path == "/v1/chat") { const std::string header = "HTTP/1.1 200 OK\r\n" "Content-Type: application/x-ndjson; charset=utf-8\r\n" "Transfer-Encoding: chunked\r\n" "Cache-Control: no-cache\r\n" "Connection: close\r\n\r\n"; if (!SendAll(client.get(), header)) return; try { engine.Chat(request.body, MaxTokens(request.target), [&](std::string_view event) { return SendChunk(client.get(), event); }); } catch (const std::exception & error) { const std::string event = "{\"type\":\"error\",\"message\":\"" + JsonEscape(error.what()) + "\"}\n"; SendChunk(client.get(), event); } SendAll(client.get(), "0\r\n\r\n"); return; } SendResponse(client.get(), 404, "Not Found", "{\"error\":\"not found\"}\n"); } catch (const std::exception & error) { const std::string body = "{\"error\":\"" + JsonEscape(error.what()) + "\"}\n"; SendResponse(client.get(), 400, "Bad Request", body); } } #endif } // namespace int RunHttpService( const std::filesystem::path & package_path, std::uint16_t port, std::string_view bind_address) { #if !defined(__linux__) (void)package_path; (void)port; (void)bind_address; throw std::runtime_error("the HTTP service currently requires Linux"); #else std::signal(SIGPIPE, SIG_IGN); const auto initialize_begin = Clock::now(); const ModelPackage package(package_path); const auto tokenizer = PackageTokenizer(package); Decoder decoder(package); const std::size_t batch_granularity = decoder.batch_granularity(); if (batch_granularity != 0) { for (std::size_t rows = batch_granularity; rows <= 128; rows *= 2) { decoder.PrepareBatch(rows); } } const double initialization_ms = Milliseconds(initialize_begin, Clock::now()); SessionEngine engine(package, tokenizer, decoder, initialization_ms); Socket server(socket(AF_INET, SOCK_STREAM | SOCK_CLOEXEC, 0)); if (server.get() < 0) { throw std::runtime_error(std::string("socket: ") + std::strerror(errno)); } int reuse = 1; setsockopt(server.get(), SOL_SOCKET, SO_REUSEADDR, &reuse, sizeof(reuse)); sockaddr_in address {}; address.sin_family = AF_INET; address.sin_port = htons(port); const std::string bind_text(bind_address); if (inet_pton(AF_INET, bind_text.c_str(), &address.sin_addr) != 1) { throw std::invalid_argument("bind address must be an IPv4 address"); } if (bind(server.get(), reinterpret_cast(&address), sizeof(address)) != 0) { throw std::runtime_error(std::string("bind: ") + std::strerror(errno)); } if (listen(server.get(), 16) != 0) { throw std::runtime_error(std::string("listen: ") + std::strerror(errno)); } std::cout << std::fixed << std::setprecision(3) << "service_ready=http://" << bind_address << ':' << port << '\n' << "initialization_ms=" << initialization_ms << '\n' << "prewarm_experts=" << (std::getenv("LING3_PREWARM_EXPERTS") == nullptr ? "false" : "true") << '\n' << "dynamic_batch_1_128=" << (decoder.has_dynamic_batch() ? "true" : "false") << std::endl; while (true) { const int descriptor = accept4(server.get(), nullptr, nullptr, SOCK_CLOEXEC); if (descriptor < 0 && errno == EINTR) continue; if (descriptor < 0) { throw std::runtime_error(std::string("accept: ") + std::strerror(errno)); } std::thread([descriptor, &engine]() { HandleClient(Socket(descriptor), engine); }).detach(); } #endif } } // namespace ling3