KBaba7/llama.cpp
0
1/*2 Copyright 2024 Google LLC3 4 Use of this source code is governed by an MIT-style5 license that can be found in the LICENSE file or at6 https://opensource.org/licenses/MIT.7*/8// SPDX-License-Identifier: MIT9#pragma once10 11#include <iostream>12#include <string>13#include <vector>14#include <regex>15#include <memory>16#include <stdexcept>17#include <sstream>18#include <unordered_set>19#include <json.hpp>20 21using json = nlohmann::ordered_json;22 23namespace minja {24 25class Context;26 27struct Options {28 bool trim_blocks; // removes the first newline after a block29 bool lstrip_blocks; // removes leading whitespace on the line of the block30 bool keep_trailing_newline; // don't remove last newline31};32 33struct ArgumentsValue;34 35inline std::string normalize_newlines(const std::string & s) {36#ifdef _WIN3237 static const std::regex nl_regex("\r\n");38 return std::regex_replace(s, nl_regex, "\n");39#else40 return s;41#endif42}43 44/* Values that behave roughly like in Python. */45class Value : public std::enable_shared_from_this<Value> {46public:47 using CallableType = std::function<Value(const std::shared_ptr<Context> &, ArgumentsValue &)>;48 using FilterType = std::function<Value(const std::shared_ptr<Context> &, ArgumentsValue &)>;49 50private:51 using ObjectType = nlohmann::ordered_map<json, Value>; // Only contains primitive keys52 using ArrayType = std::vector<Value>;53 54 std::shared_ptr<ArrayType> array_;55 std::shared_ptr<ObjectType> object_;56 std::shared_ptr<CallableType> callable_;57 json primitive_;58 59 Value(const std::shared_ptr<ArrayType> & array) : array_(array) {}60 Value(const std::shared_ptr<ObjectType> & object) : object_(object) {}61 Value(const std::shared_ptr<CallableType> & callable) : object_(std::make_shared<ObjectType>()), callable_(callable) {}62 63 /* Python-style string repr */64 static void dump_string(const json & primitive, std::ostringstream & out, char string_quote = '\'') {65 if (!primitive.is_string()) throw std::runtime_error("Value is not a string: " + primitive.dump());66 auto s = primitive.dump();67 if (string_quote == '"' || s.find('\'') != std::string::npos) {68 out << s;69 return;70 }71 // Reuse json dump, just changing string quotes72 out << string_quote;73 for (size_t i = 1, n = s.size() - 1; i < n; ++i) {74 if (s[i] == '\\' && s[i + 1] == '"') {75 out << '"';76 i++;77 } else if (s[i] == string_quote) {78 out << '\\' << string_quote;79 } else {80 out << s[i];81 }82 }83 out << string_quote;84 }85 void dump(std::ostringstream & out, int indent = -1, int level = 0, bool to_json = false) const {86 auto print_indent = [&](int level) {87 if (indent > 0) {88 out << "\n";89 for (int i = 0, n = level * indent; i < n; ++i) out << ' ';90 }91 };92 auto print_sub_sep = [&]() {93 out << ',';94 if (indent < 0) out << ' ';95 else print_indent(level + 1);96 };97 98 auto string_quote = to_json ? '"' : '\'';99 100 if (is_null()) out << "null";101 else if (array_) {102 out << "[";103 print_indent(level + 1);104 for (size_t i = 0; i < array_->size(); ++i) {105 if (i) print_sub_sep();106 (*array_)[i].dump(out, indent, level + 1, to_json);107 }108 print_indent(level);109 out << "]";110 } else if (object_) {111 out << "{";112 print_indent(level + 1);113 for (auto begin = object_->begin(), it = begin; it != object_->end(); ++it) {114 if (it != begin) print_sub_sep();115 if (it->first.is_string()) {116 dump_string(it->first, out, string_quote);117 } else {118 out << string_quote << it->first.dump() << string_quote;119 }120 out << ": ";121 it->second.dump(out, indent, level + 1, to_json);122 }123 print_indent(level);124 out << "}";125 } else if (callable_) {126 throw std::runtime_error("Cannot dump callable to JSON");127 } else if (is_boolean() && !to_json) {128 out << (this->to_bool() ? "True" : "False");129 } else if (is_string() && !to_json) {130 dump_string(primitive_, out, string_quote);131 } else {132 out << primitive_.dump();133 }134 }135 136public:137 Value() {}138 Value(const bool& v) : primitive_(v) {}139 Value(const int64_t & v) : primitive_(v) {}140 Value(const double& v) : primitive_(v) {}141 Value(const std::nullptr_t &) {}142 Value(const std::string & v) : primitive_(v) {}143 Value(const char * v) : primitive_(std::string(v)) {}144 145 Value(const json & v) {146 if (v.is_object()) {147 auto object = std::make_shared<ObjectType>();148 for (auto it = v.begin(); it != v.end(); ++it) {149 (*object)[it.key()] = it.value();150 }151 object_ = std::move(object);152 } else if (v.is_array()) {153 auto array = std::make_shared<ArrayType>();154 for (const auto& item : v) {155 array->push_back(Value(item));156 }157 array_ = array;158 } else {159 primitive_ = v;160 }161 }162 163 std::vector<Value> keys() {164 if (!object_) throw std::runtime_error("Value is not an object: " + dump());165 std::vector<Value> res;166 for (const auto& item : *object_) {167 res.push_back(item.first);168 }169 return res;170 }171 172 size_t size() const {173 if (is_object()) return object_->size();174 if (is_array()) return array_->size();175 if (is_string()) return primitive_.get<std::string>().length();176 throw std::runtime_error("Value is not an array or object: " + dump());177 }178 179 static Value array(const std::vector<Value> values = {}) {180 auto array = std::make_shared<ArrayType>();181 for (const auto& item : values) {182 array->push_back(item);183 }184 return Value(array);185 }186 static Value object(const std::shared_ptr<ObjectType> object = std::make_shared<ObjectType>()) {187 return Value(object);188 }189 static Value callable(const CallableType & callable) {190 return Value(std::make_shared<CallableType>(callable));191 }192 193 void insert(size_t index, const Value& v) {194 if (!array_)195 throw std::runtime_error("Value is not an array: " + dump());196 array_->insert(array_->begin() + index, v);197 }198 void push_back(const Value& v) {199 if (!array_)200 throw std::runtime_error("Value is not an array: " + dump());201 array_->push_back(v);202 }203 Value pop(const Value& index) {204 if (is_array()) {205 if (array_->empty())206 throw std::runtime_error("pop from empty list");207 if (index.is_null()) {208 auto ret = array_->back();209 array_->pop_back();210 return ret;211 } else if (!index.is_number_integer()) {212 throw std::runtime_error("pop index must be an integer: " + index.dump());213 } else {214 auto i = index.get<int>();215 if (i < 0 || i >= static_cast<int>(array_->size()))216 throw std::runtime_error("pop index out of range: " + index.dump());217 auto it = array_->begin() + (i < 0 ? array_->size() + i : i);218 auto ret = *it;219 array_->erase(it);220 return ret;221 }222 } else if (is_object()) {223 if (!index.is_hashable())224 throw std::runtime_error("Unashable type: " + index.dump());225 auto it = object_->find(index.primitive_);226 if (it == object_->end())227 throw std::runtime_error("Key not found: " + index.dump());228 auto ret = it->second;229 object_->erase(it);230 return ret;231 } else {232 throw std::runtime_error("Value is not an array or object: " + dump());233 }234 }235 Value get(const Value& key) {236 if (array_) {237 if (!key.is_number_integer()) {238 return Value();239 }240 auto index = key.get<int>();241 return array_->at(index < 0 ? array_->size() + index : index);242 } else if (object_) {243 if (!key.is_hashable()) throw std::runtime_error("Unashable type: " + dump());244 auto it = object_->find(key.primitive_);245 if (it == object_->end()) return Value();246 return it->second;247 }248 return Value();249 }250 void set(const Value& key, const Value& value) {251 if (!object_) throw std::runtime_error("Value is not an object: " + dump());252 if (!key.is_hashable()) throw std::runtime_error("Unashable type: " + dump());253 (*object_)[key.primitive_] = value;254 }255 Value call(const std::shared_ptr<Context> & context, ArgumentsValue & args) const {256 if (!callable_) throw std::runtime_error("Value is not callable: " + dump());257 return (*callable_)(context, args);258 }259 260 bool is_object() const { return !!object_; }261 bool is_array() const { return !!array_; }262 bool is_callable() const { return !!callable_; }263 bool is_null() const { return !object_ && !array_ && primitive_.is_null() && !callable_; }264 bool is_boolean() const { return primitive_.is_boolean(); }265 bool is_number_integer() const { return primitive_.is_number_integer(); }266 bool is_number_float() const { return primitive_.is_number_float(); }267 bool is_number() const { return primitive_.is_number(); }268 bool is_string() const { return primitive_.is_string(); }269 bool is_iterable() const { return is_array() || is_object() || is_string(); }270 271 bool is_primitive() const { return !array_ && !object_ && !callable_; }272 bool is_hashable() const { return is_primitive(); }273 274 bool empty() const {275 if (is_null())276 throw std::runtime_error("Undefined value or reference");277 if (is_string()) return primitive_.empty();278 if (is_array()) return array_->empty();279 if (is_object()) return object_->empty();280 return false;281 }282 283 void for_each(const std::function<void(Value &)> & callback) const {284 if (is_null())285 throw std::runtime_error("Undefined value or reference");286 if (array_) {287 for (auto& item : *array_) {288 callback(item);289 }290 } else if (object_) {291 for (auto & item : *object_) {292 Value key(item.first);293 callback(key);294 }295 } else if (is_string()) {296 for (char c : primitive_.get<std::string>()) {297 auto val = Value(std::string(1, c));298 callback(val);299 }300 } else {301 throw std::runtime_error("Value is not iterable: " + dump());302 }303 }304 305 bool to_bool() const {306 if (is_null()) return false;307 if (is_boolean()) return get<bool>();308 if (is_number()) return get<double>() != 0;309 if (is_string()) return !get<std::string>().empty();310 if (is_array()) return !empty();311 return true;312 }313 314 int64_t to_int() const {315 if (is_null()) return 0;316 if (is_boolean()) return get<bool>() ? 1 : 0;317 if (is_number()) return static_cast<int64_t>(get<double>());318 if (is_string()) {319 try {320 return std::stol(get<std::string>());321 } catch (const std::exception &) {322 return 0;323 }324 }325 return 0;326 }327 328 bool operator<(const Value & other) const {329 if (is_null())330 throw std::runtime_error("Undefined value or reference");331 if (is_number() && other.is_number()) return get<double>() < other.get<double>();332 if (is_string() && other.is_string()) return get<std::string>() < other.get<std::string>();333 throw std::runtime_error("Cannot compare values: " + dump() + " < " + other.dump());334 }335 bool operator>=(const Value & other) const { return !(*this < other); }336 337 bool operator>(const Value & other) const {338 if (is_null())339 throw std::runtime_error("Undefined value or reference");340 if (is_number() && other.is_number()) return get<double>() > other.get<double>();341 if (is_string() && other.is_string()) return get<std::string>() > other.get<std::string>();342 throw std::runtime_error("Cannot compare values: " + dump() + " > " + other.dump());343 }344 bool operator<=(const Value & other) const { return !(*this > other); }345 346 bool operator==(const Value & other) const {347 if (callable_ || other.callable_) {348 if (callable_.get() != other.callable_.get()) return false;349 }350 if (array_) {351 if (!other.array_) return false;352 if (array_->size() != other.array_->size()) return false;353 for (size_t i = 0; i < array_->size(); ++i) {354 if (!(*array_)[i].to_bool() || !(*other.array_)[i].to_bool() || (*array_)[i] != (*other.array_)[i]) return false;355 }356 return true;357 } else if (object_) {358 if (!other.object_) return false;359 if (object_->size() != other.object_->size()) return false;360 for (const auto& item : *object_) {361 if (!item.second.to_bool() || !other.object_->count(item.first) || item.second != other.object_->at(item.first)) return false;362 }363 return true;364 } else {365 return primitive_ == other.primitive_;366 }367 }368 bool operator!=(const Value & other) const { return !(*this == other); }369 370 bool contains(const char * key) const { return contains(std::string(key)); }371 bool contains(const std::string & key) const {372 if (array_) {373 return false;374 } else if (object_) {375 return object_->find(key) != object_->end();376 } else {377 throw std::runtime_error("contains can only be called on arrays and objects: " + dump());378 }379 }380 bool contains(const Value & value) const {381 if (is_null())382 throw std::runtime_error("Undefined value or reference");383 if (array_) {384 for (const auto& item : *array_) {385 if (item.to_bool() && item == value) return true;386 }387 return false;388 } else if (object_) {389 if (!value.is_hashable()) throw std::runtime_error("Unashable type: " + value.dump());390 return object_->find(value.primitive_) != object_->end();391 } else {392 throw std::runtime_error("contains can only be called on arrays and objects: " + dump());393 }394 }395 void erase(size_t index) {396 if (!array_) throw std::runtime_error("Value is not an array: " + dump());397 array_->erase(array_->begin() + index);398 }399 void erase(const std::string & key) {400 if (!object_) throw std::runtime_error("Value is not an object: " + dump());401 object_->erase(key);402 }403 const Value& at(const Value & index) const {404 return const_cast<Value*>(this)->at(index);405 }406 Value& at(const Value & index) {407 if (!index.is_hashable()) throw std::runtime_error("Unashable type: " + dump());408 if (is_array()) return array_->at(index.get<int>());409 if (is_object()) return object_->at(index.primitive_);410 throw std::runtime_error("Value is not an array or object: " + dump());411 }412 const Value& at(size_t index) const {413 return const_cast<Value*>(this)->at(index);414 }415 Value& at(size_t index) {416 if (is_null())417 throw std::runtime_error("Undefined value or reference");418 if (is_array()) return array_->at(index);419 if (is_object()) return object_->at(index);420 throw std::runtime_error("Value is not an array or object: " + dump());421 }422 423 template <typename T>424 T get(const std::string & key, T default_value) const {425 if (!contains(key)) return default_value;426 return at(key).get<T>();427 }428 429 template <typename T>430 T get() const {431 if (is_primitive()) return primitive_.get<T>();432 throw std::runtime_error("get<T> not defined for this value type: " + dump());433 }434 435 std::string dump(int indent=-1, bool to_json=false) const {436 std::ostringstream out;437 dump(out, indent, 0, to_json);438 return out.str();439 }440 441 Value operator-() const {442 if (is_number_integer())443 return -get<int64_t>();444 else445 return -get<double>();446 }447 std::string to_str() const {448 if (is_string()) return get<std::string>();449 if (is_number_integer()) return std::to_string(get<int64_t>());450 if (is_number_float()) return std::to_string(get<double>());451 if (is_boolean()) return get<bool>() ? "True" : "False";452 if (is_null()) return "None";453 return dump();454 }455 Value operator+(const Value& rhs) const {456 if (is_string() || rhs.is_string()) {457 return to_str() + rhs.to_str();458 } else if (is_number_integer() && rhs.is_number_integer()) {459 return get<int64_t>() + rhs.get<int64_t>();460 } else if (is_array() && rhs.is_array()) {461 auto res = Value::array();462 for (const auto& item : *array_) res.push_back(item);463 for (const auto& item : *rhs.array_) res.push_back(item);464 return res;465 } else {466 return get<double>() + rhs.get<double>();467 }468 }469 Value operator-(const Value& rhs) const {470 if (is_number_integer() && rhs.is_number_integer())471 return get<int64_t>() - rhs.get<int64_t>();472 else473 return get<double>() - rhs.get<double>();474 }475 Value operator*(const Value& rhs) const {476 if (is_string() && rhs.is_number_integer()) {477 std::ostringstream out;478 for (int64_t i = 0, n = rhs.get<int64_t>(); i < n; ++i) {479 out << to_str();480 }481 return out.str();482 }483 else if (is_number_integer() && rhs.is_number_integer())484 return get<int64_t>() * rhs.get<int64_t>();485 else486 return get<double>() * rhs.get<double>();487 }488 Value operator/(const Value& rhs) const {489 if (is_number_integer() && rhs.is_number_integer())490 return get<int64_t>() / rhs.get<int64_t>();491 else492 return get<double>() / rhs.get<double>();493 }494 Value operator%(const Value& rhs) const {495 return get<int64_t>() % rhs.get<int64_t>();496 }497};498 499struct ArgumentsValue {500 std::vector<Value> args;501 std::vector<std::pair<std::string, Value>> kwargs;502 503 bool has_named(const std::string & name) {504 for (const auto & p : kwargs) {505 if (p.first == name) return true;506 }507 return false;508 }509 510 Value get_named(const std::string & name) {511 for (const auto & [key, value] : kwargs) {512 if (key == name) return value;513 }514 return Value();515 }516 517 bool empty() {518 return args.empty() && kwargs.empty();519 }520 521 void expectArgs(const std::string & method_name, const std::pair<size_t, size_t> & pos_count, const std::pair<size_t, size_t> & kw_count) {522 if (args.size() < pos_count.first || args.size() > pos_count.second || kwargs.size() < kw_count.first || kwargs.size() > kw_count.second) {523 std::ostringstream out;524 out << method_name << " must have between " << pos_count.first << " and " << pos_count.second << " positional arguments and between " << kw_count.first << " and " << kw_count.second << " keyword arguments";525 throw std::runtime_error(out.str());526 }527 }528};529 530template <>531inline json Value::get<json>() const {532 if (is_primitive()) return primitive_;533 if (is_null()) return json();534 if (array_) {535 std::vector<json> res;536 for (const auto& item : *array_) {537 res.push_back(item.get<json>());538 }539 return res;540 }541 if (object_) {542 json res = json::object();543 for (const auto& [key, value] : *object_) {544 if (key.is_string()) {545 res[key.get<std::string>()] = value.get<json>();546 } else if (key.is_primitive()) {547 res[key.dump()] = value.get<json>();548 } else {549 throw std::runtime_error("Invalid key type for conversion to JSON: " + key.dump());550 }551 }552 if (is_callable()) {553 res["__callable__"] = true;554 }555 return res;556 }557 throw std::runtime_error("get<json> not defined for this value type: " + dump());558}559 560} // namespace minja561 562namespace std {563 template <>564 struct hash<minja::Value> {565 size_t operator()(const minja::Value & v) const {566 if (!v.is_hashable())567 throw std::runtime_error("Unsupported type for hashing: " + v.dump());568 return std::hash<json>()(v.get<json>());569 }570 };571} // namespace std572 573namespace minja {574 575static std::string error_location_suffix(const std::string & source, size_t pos) {576 auto get_line = [&](size_t line) {577 auto start = source.begin();578 for (size_t i = 1; i < line; ++i) {579 start = std::find(start, source.end(), '\n') + 1;580 }581 auto end = std::find(start, source.end(), '\n');582 return std::string(start, end);583 };584 auto start = source.begin();585 auto end = source.end();586 auto it = start + pos;587 auto line = std::count(start, it, '\n') + 1;588 auto max_line = std::count(start, end, '\n') + 1;589 auto col = pos - std::string(start, it).rfind('\n');590 std::ostringstream out;591 out << " at row " << line << ", column " << col << ":\n";592 if (line > 1) out << get_line(line - 1) << "\n";593 out << get_line(line) << "\n";594 out << std::string(col - 1, ' ') << "^\n";595 if (line < max_line) out << get_line(line + 1) << "\n";596 597 return out.str();598}599 600class Context : public std::enable_shared_from_this<Context> {601 protected:602 Value values_;603 std::shared_ptr<Context> parent_;604 public:605 Context(Value && values, const std::shared_ptr<Context> & parent = nullptr) : values_(std::move(values)), parent_(parent) {606 if (!values_.is_object()) throw std::runtime_error("Context values must be an object: " + values_.dump());607 }608 virtual ~Context() {}609 610 static std::shared_ptr<Context> builtins();611 static std::shared_ptr<Context> make(Value && values, const std::shared_ptr<Context> & parent = builtins());612 613 std::vector<Value> keys() {614 return values_.keys();615 }616 virtual Value get(const Value & key) {617 if (values_.contains(key)) return values_.at(key);618 if (parent_) return parent_->get(key);619 return Value();620 }621 virtual Value & at(const Value & key) {622 if (values_.contains(key)) return values_.at(key);623 if (parent_) return parent_->at(key);624 throw std::runtime_error("Undefined variable: " + key.dump());625 }626 virtual bool contains(const Value & key) {627 if (values_.contains(key)) return true;628 if (parent_) return parent_->contains(key);629 return false;630 }631 virtual void set(const Value & key, const Value & value) {632 values_.set(key, value);633 }634};635 636struct Location {637 std::shared_ptr<std::string> source;638 size_t pos;639};640 641class Expression {642protected:643 virtual Value do_evaluate(const std::shared_ptr<Context> & context) const = 0;644public:645 using Parameters = std::vector<std::pair<std::string, std::shared_ptr<Expression>>>;646 647 Location location;648 649 Expression(const Location & location) : location(location) {}650 virtual ~Expression() = default;651 652 Value evaluate(const std::shared_ptr<Context> & context) const {653 try {654 return do_evaluate(context);655 } catch (const std::exception & e) {656 std::ostringstream out;657 out << e.what();658 if (location.source) out << error_location_suffix(*location.source, location.pos);659 throw std::runtime_error(out.str());660 }661 }662};663 664class VariableExpr : public Expression {665 std::string name;666public:667 VariableExpr(const Location & location, const std::string& n)668 : Expression(location), name(n) {}669 std::string get_name() const { return name; }670 Value do_evaluate(const std::shared_ptr<Context> & context) const override {671 if (!context->contains(name)) {672 return Value();673 }674 return context->at(name);675 }676};677 678static void destructuring_assign(const std::vector<std::string> & var_names, const std::shared_ptr<Context> & context, Value& item) {679 if (var_names.size() == 1) {680 Value name(var_names[0]);681 context->set(name, item);682 } else {683 if (!item.is_array() || item.size() != var_names.size()) {684 throw std::runtime_error("Mismatched number of variables and items in destructuring assignment");685 }686 for (size_t i = 0; i < var_names.size(); ++i) {687 context->set(var_names[i], item.at(i));688 }689 }690}691 692enum SpaceHandling { Keep, Strip, StripSpaces, StripNewline };693 694class TemplateToken {695public:696 enum class Type { Text, Expression, If, Else, Elif, EndIf, For, EndFor, Generation, EndGeneration, Set, EndSet, Comment, Macro, EndMacro, Filter, EndFilter, Break, Continue };697 698 static std::string typeToString(Type t) {699 switch (t) {700 case Type::Text: return "text";701 case Type::Expression: return "expression";702 case Type::If: return "if";703 case Type::Else: return "else";704 case Type::Elif: return "elif";705 case Type::EndIf: return "endif";706 case Type::For: return "for";707 case Type::EndFor: return "endfor";708 case Type::Set: return "set";709 case Type::EndSet: return "endset";710 case Type::Comment: return "comment";711 case Type::Macro: return "macro";712 case Type::EndMacro: return "endmacro";713 case Type::Filter: return "filter";714 case Type::EndFilter: return "endfilter";715 case Type::Generation: return "generation";716 case Type::EndGeneration: return "endgeneration";717 case Type::Break: return "break";718 case Type::Continue: return "continue";719 }720 return "Unknown";721 }722 723 TemplateToken(Type type, const Location & location, SpaceHandling pre, SpaceHandling post) : type(type), location(location), pre_space(pre), post_space(post) {}724 virtual ~TemplateToken() = default;725 726 Type type;727 Location location;728 SpaceHandling pre_space = SpaceHandling::Keep;729 SpaceHandling post_space = SpaceHandling::Keep;730};731 732struct TextTemplateToken : public TemplateToken {733 std::string text;734 TextTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, const std::string& t) : TemplateToken(Type::Text, location, pre, post), text(t) {}735};736 737struct ExpressionTemplateToken : public TemplateToken {738 std::shared_ptr<Expression> expr;739 ExpressionTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, std::shared_ptr<Expression> && e) : TemplateToken(Type::Expression, location, pre, post), expr(std::move(e)) {}740};741 742struct IfTemplateToken : public TemplateToken {743 std::shared_ptr<Expression> condition;744 IfTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, std::shared_ptr<Expression> && c) : TemplateToken(Type::If, location, pre, post), condition(std::move(c)) {}745};746 747struct ElifTemplateToken : public TemplateToken {748 std::shared_ptr<Expression> condition;749 ElifTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, std::shared_ptr<Expression> && c) : TemplateToken(Type::Elif, location, pre, post), condition(std::move(c)) {}750};751 752struct ElseTemplateToken : public TemplateToken {753 ElseTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::Else, location, pre, post) {}754};755 756struct EndIfTemplateToken : public TemplateToken {757 EndIfTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::EndIf, location, pre, post) {}758};759 760struct MacroTemplateToken : public TemplateToken {761 std::shared_ptr<VariableExpr> name;762 Expression::Parameters params;763 MacroTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, std::shared_ptr<VariableExpr> && n, Expression::Parameters && p)764 : TemplateToken(Type::Macro, location, pre, post), name(std::move(n)), params(std::move(p)) {}765};766 767struct EndMacroTemplateToken : public TemplateToken {768 EndMacroTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::EndMacro, location, pre, post) {}769};770 771struct FilterTemplateToken : public TemplateToken {772 std::shared_ptr<Expression> filter;773 FilterTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, std::shared_ptr<Expression> && filter)774 : TemplateToken(Type::Filter, location, pre, post), filter(std::move(filter)) {}775};776 777struct EndFilterTemplateToken : public TemplateToken {778 EndFilterTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::EndFilter, location, pre, post) {}779};780 781struct ForTemplateToken : public TemplateToken {782 std::vector<std::string> var_names;783 std::shared_ptr<Expression> iterable;784 std::shared_ptr<Expression> condition;785 bool recursive;786 ForTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, const std::vector<std::string> & vns, std::shared_ptr<Expression> && iter,787 std::shared_ptr<Expression> && c, bool r)788 : TemplateToken(Type::For, location, pre, post), var_names(vns), iterable(std::move(iter)), condition(std::move(c)), recursive(r) {}789};790 791struct EndForTemplateToken : public TemplateToken {792 EndForTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::EndFor, location, pre, post) {}793};794 795struct GenerationTemplateToken : public TemplateToken {796 GenerationTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::Generation, location, pre, post) {}797};798 799struct EndGenerationTemplateToken : public TemplateToken {800 EndGenerationTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::EndGeneration, location, pre, post) {}801};802 803struct SetTemplateToken : public TemplateToken {804 std::string ns;805 std::vector<std::string> var_names;806 std::shared_ptr<Expression> value;807 SetTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, const std::string & ns, const std::vector<std::string> & vns, std::shared_ptr<Expression> && v)808 : TemplateToken(Type::Set, location, pre, post), ns(ns), var_names(vns), value(std::move(v)) {}809};810 811struct EndSetTemplateToken : public TemplateToken {812 EndSetTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post) : TemplateToken(Type::EndSet, location, pre, post) {}813};814 815struct CommentTemplateToken : public TemplateToken {816 std::string text;817 CommentTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, const std::string& t) : TemplateToken(Type::Comment, location, pre, post), text(t) {}818};819 820enum class LoopControlType { Break, Continue };821 822class LoopControlException : public std::runtime_error {823public:824 LoopControlType control_type;825 LoopControlException(const std::string & message, LoopControlType control_type) : std::runtime_error(message), control_type(control_type) {}826 LoopControlException(LoopControlType control_type)827 : std::runtime_error((control_type == LoopControlType::Continue ? "continue" : "break") + std::string(" outside of a loop")),828 control_type(control_type) {}829};830 831struct LoopControlTemplateToken : public TemplateToken {832 LoopControlType control_type;833 LoopControlTemplateToken(const Location & location, SpaceHandling pre, SpaceHandling post, LoopControlType control_type) : TemplateToken(Type::Break, location, pre, post), control_type(control_type) {}834};835 836class TemplateNode {837 Location location_;838protected:839 virtual void do_render(std::ostringstream & out, const std::shared_ptr<Context> & context) const = 0;840 841public:842 TemplateNode(const Location & location) : location_(location) {}843 void render(std::ostringstream & out, const std::shared_ptr<Context> & context) const {844 try {845 do_render(out, context);846 } catch (const LoopControlException & e) {847 // TODO: make stack creation lazy. Only needed if it was thrown outside of a loop.848 std::ostringstream err;849 err << e.what();850 if (location_.source) err << error_location_suffix(*location_.source, location_.pos);851 throw LoopControlException(err.str(), e.control_type);852 } catch (const std::exception & e) {853 std::ostringstream err;854 err << e.what();855 if (location_.source) err << error_location_suffix(*location_.source, location_.pos);856 throw std::runtime_error(err.str());857 }858 }859 const Location & location() const { return location_; }860 virtual ~TemplateNode() = default;861 std::string render(const std::shared_ptr<Context> & context) const {862 std::ostringstream out;863 render(out, context);864 return out.str();865 }866};867 868class SequenceNode : public TemplateNode {869 std::vector<std::shared_ptr<TemplateNode>> children;870public:871 SequenceNode(const Location & location, std::vector<std::shared_ptr<TemplateNode>> && c)872 : TemplateNode(location), children(std::move(c)) {}873 void do_render(std::ostringstream & out, const std::shared_ptr<Context> & context) const override {874 for (const auto& child : children) child->render(out, context);875 }876};877 878class TextNode : public TemplateNode {879 std::string text;880public:881 TextNode(const Location & location, const std::string& t) : TemplateNode(location), text(t) {}882 void do_render(std::ostringstream & out, const std::shared_ptr<Context> &) const override {883 out << text;884 }885};886 887class ExpressionNode : public TemplateNode {888 std::shared_ptr<Expression> expr;889public:890 ExpressionNode(const Location & location, std::shared_ptr<Expression> && e) : TemplateNode(location), expr(std::move(e)) {}891 void do_render(std::ostringstream & out, const std::shared_ptr<Context> & context) const override {892 if (!expr) throw std::runtime_error("ExpressionNode.expr is null");893 auto result = expr->evaluate(context);894 if (result.is_string()) {895 out << result.get<std::string>();896 } else if (result.is_boolean()) {897 out << (result.get<bool>() ? "True" : "False");898 } else if (!result.is_null()) {899 out << result.dump();900 }901 }902};903 904class IfNode : public TemplateNode {905 std::vector<std::pair<std::shared_ptr<Expression>, std::shared_ptr<TemplateNode>>> cascade;906public:907 IfNode(const Location & location, std::vector<std::pair<std::shared_ptr<Expression>, std::shared_ptr<TemplateNode>>> && c)908 : TemplateNode(location), cascade(std::move(c)) {}909 void do_render(std::ostringstream & out, const std::shared_ptr<Context> & context) const override {910 for (const auto& branch : cascade) {911 auto enter_branch = true;912 if (branch.first) {913 enter_branch = branch.first->evaluate(context).to_bool();914 }915 if (enter_branch) {916 if (!branch.second) throw std::runtime_error("IfNode.cascade.second is null");917 branch.second->render(out, context);918 return;919 }920 }921 }922};923 924class LoopControlNode : public TemplateNode {925 LoopControlType control_type_;926 public:927 LoopControlNode(const Location & location, LoopControlType control_type) : TemplateNode(location), control_type_(control_type) {}928 void do_render(std::ostringstream &, const std::shared_ptr<Context> &) const override {929 throw LoopControlException(control_type_);930 }931};932 933class ForNode : public TemplateNode {934 std::vector<std::string> var_names;935 std::shared_ptr<Expression> iterable;936 std::shared_ptr<Expression> condition;937 std::shared_ptr<TemplateNode> body;938 bool recursive;939 std::shared_ptr<TemplateNode> else_body;940public:941 ForNode(const Location & location, std::vector<std::string> && var_names, std::shared_ptr<Expression> && iterable,942 std::shared_ptr<Expression> && condition, std::shared_ptr<TemplateNode> && body, bool recursive, std::shared_ptr<TemplateNode> && else_body)943 : TemplateNode(location), var_names(var_names), iterable(std::move(iterable)), condition(std::move(condition)), body(std::move(body)), recursive(recursive), else_body(std::move(else_body)) {}944 945 void do_render(std::ostringstream & out, const std::shared_ptr<Context> & context) const override {946 // https://jinja.palletsprojects.com/en/3.0.x/templates/#for947 if (!iterable) throw std::runtime_error("ForNode.iterable is null");948 if (!body) throw std::runtime_error("ForNode.body is null");949 950 auto iterable_value = iterable->evaluate(context);951 Value::CallableType loop_function;952 953 std::function<void(Value&)> visit = [&](Value& iter) {954 auto filtered_items = Value::array();955 if (!iter.is_null()) {956 if (!iterable_value.is_iterable()) {957 throw std::runtime_error("For loop iterable must be iterable: " + iterable_value.dump());958 }959 iterable_value.for_each([&](Value & item) {960 destructuring_assign(var_names, context, item);961 if (!condition || condition->evaluate(context).to_bool()) {962 filtered_items.push_back(item);963 }964 });965 }966 if (filtered_items.empty()) {967 if (else_body) {968 else_body->render(out, context);969 }970 } else {971 auto loop = recursive ? Value::callable(loop_function) : Value::object();972 loop.set("length", (int64_t) filtered_items.size());973 974 size_t cycle_index = 0;975 loop.set("cycle", Value::callable([&](const std::shared_ptr<Context> &, ArgumentsValue & args) {976 if (args.args.empty() || !args.kwargs.empty()) {977 throw std::runtime_error("cycle() expects at least 1 positional argument and no named arg");978 }979 auto item = args.args[cycle_index];980 cycle_index = (cycle_index + 1) % args.args.size();981 return item;982 }));983 auto loop_context = Context::make(Value::object(), context);984 loop_context->set("loop", loop);985 for (size_t i = 0, n = filtered_items.size(); i < n; ++i) {986 auto & item = filtered_items.at(i);987 destructuring_assign(var_names, loop_context, item);988 loop.set("index", (int64_t) i + 1);989 loop.set("index0", (int64_t) i);990 loop.set("revindex", (int64_t) (n - i));991 loop.set("revindex0", (int64_t) (n - i - 1));992 loop.set("length", (int64_t) n);993 loop.set("first", i == 0);994 loop.set("last", i == (n - 1));995 loop.set("previtem", i > 0 ? filtered_items.at(i - 1) : Value());996 loop.set("nextitem", i < n - 1 ? filtered_items.at(i + 1) : Value());997 try {998 body->render(out, loop_context);999 } catch (const LoopControlException & e) {1000 if (e.control_type == LoopControlType::Break) break;1001 if (e.control_type == LoopControlType::Continue) continue;1002 }1003 }1004 }1005 };1006 1007 if (recursive) {1008 loop_function = [&](const std::shared_ptr<Context> &, ArgumentsValue & args) {1009 if (args.args.size() != 1 || !args.kwargs.empty() || !args.args[0].is_array()) {1010 throw std::runtime_error("loop() expects exactly 1 positional iterable argument");1011 }1012 auto & items = args.args[0];1013 visit(items);1014 return Value();1015 };1016 }1017 1018 visit(iterable_value);1019 }1020};1021 1022class MacroNode : public TemplateNode {1023 std::shared_ptr<VariableExpr> name;1024 Expression::Parameters params;1025 std::shared_ptr<TemplateNode> body;1026 std::unordered_map<std::string, size_t> named_param_positions;1027public:1028 MacroNode(const Location & location, std::shared_ptr<VariableExpr> && n, Expression::Parameters && p, std::shared_ptr<TemplateNode> && b)1029 : TemplateNode(location), name(std::move(n)), params(std::move(p)), body(std::move(b)) {1030 for (size_t i = 0; i < params.size(); ++i) {1031 const auto & name = params[i].first;1032 if (!name.empty()) {1033 named_param_positions[name] = i;1034 }1035 }1036 }1037 void do_render(std::ostringstream &, const std::shared_ptr<Context> & macro_context) const override {1038 if (!name) throw std::runtime_error("MacroNode.name is null");1039 if (!body) throw std::runtime_error("MacroNode.body is null");1040 auto callable = Value::callable([&](const std::shared_ptr<Context> & context, ArgumentsValue & args) {1041 auto call_context = macro_context;1042 std::vector<bool> param_set(params.size(), false);1043 for (size_t i = 0, n = args.args.size(); i < n; i++) {1044 auto & arg = args.args[i];1045 if (i >= params.size()) throw std::runtime_error("Too many positional arguments for macro " + name->get_name());1046 param_set[i] = true;1047 auto & param_name = params[i].first;1048 call_context->set(param_name, arg);1049 }1050 for (auto & [arg_name, value] : args.kwargs) {1051 auto it = named_param_positions.find(arg_name);1052 if (it == named_param_positions.end()) throw std::runtime_error("Unknown parameter name for macro " + name->get_name() + ": " + arg_name);1053 1054 call_context->set(arg_name, value);1055 param_set[it->second] = true;1056 }1057 // Set default values for parameters that were not passed1058 for (size_t i = 0, n = params.size(); i < n; i++) {1059 if (!param_set[i] && params[i].second != nullptr) {1060 auto val = params[i].second->evaluate(context);1061 call_context->set(params[i].first, val);1062 }1063 }1064 return body->render(call_context);1065 });1066 macro_context->set(name->get_name(), callable);1067 }1068};1069 1070class FilterNode : public TemplateNode {1071 std::shared_ptr<Expression> filter;1072 std::shared_ptr<TemplateNode> body;1073 1074public:1075 FilterNode(const Location & location, std::shared_ptr<Expression> && f, std::shared_ptr<TemplateNode> && b)1076 : TemplateNode(location), filter(std::move(f)), body(std::move(b)) {}1077 1078 void do_render(std::ostringstream & out, const std::shared_ptr<Context> & context) const override {1079 if (!filter) throw std::runtime_error("FilterNode.filter is null");1080 if (!body) throw std::runtime_error("FilterNode.body is null");1081 auto filter_value = filter->evaluate(context);1082 if (!filter_value.is_callable()) {1083 throw std::runtime_error("Filter must be a callable: " + filter_value.dump());1084 }1085 std::string rendered_body = body->render(context);1086 1087 ArgumentsValue filter_args = {{Value(rendered_body)}, {}};1088 auto result = filter_value.call(context, filter_args);1089 out << result.to_str();1090 }1091};1092 1093class SetNode : public TemplateNode {1094 std::string ns;1095 std::vector<std::string> var_names;1096 std::shared_ptr<Expression> value;1097public:1098 SetNode(const Location & location, const std::string & ns, const std::vector<std::string> & vns, std::shared_ptr<Expression> && v)1099 : TemplateNode(location), ns(ns), var_names(vns), value(std::move(v)) {}1100 void do_render(std::ostringstream &, const std::shared_ptr<Context> & context) const override {1101 if (!value) throw std::runtime_error("SetNode.value is null");1102 if (!ns.empty()) {1103 if (var_names.size() != 1) {1104 throw std::runtime_error("Namespaced set only supports a single variable name");1105 }1106 auto & name = var_names[0];1107 auto ns_value = context->get(ns);1108 if (!ns_value.is_object()) throw std::runtime_error("Namespace '" + ns + "' is not an object");1109 ns_value.set(name, this->value->evaluate(context));1110 } else {1111 auto val = value->evaluate(context);1112 destructuring_assign(var_names, context, val);1113 }1114 }1115};1116 1117class SetTemplateNode : public TemplateNode {1118 std::string name;1119 std::shared_ptr<TemplateNode> template_value;1120public:1121 SetTemplateNode(const Location & location, const std::string & name, std::shared_ptr<TemplateNode> && tv)1122 : TemplateNode(location), name(name), template_value(std::move(tv)) {}1123 void do_render(std::ostringstream &, const std::shared_ptr<Context> & context) const override {1124 if (!template_value) throw std::runtime_error("SetTemplateNode.template_value is null");1125 Value value { template_value->render(context) };1126 context->set(name, value);1127 }1128};1129 1130class IfExpr : public Expression {1131 std::shared_ptr<Expression> condition;1132 std::shared_ptr<Expression> then_expr;1133 std::shared_ptr<Expression> else_expr;1134public:1135 IfExpr(const Location & location, std::shared_ptr<Expression> && c, std::shared_ptr<Expression> && t, std::shared_ptr<Expression> && e)1136 : Expression(location), condition(std::move(c)), then_expr(std::move(t)), else_expr(std::move(e)) {}1137 Value do_evaluate(const std::shared_ptr<Context> & context) const override {1138 if (!condition) throw std::runtime_error("IfExpr.condition is null");1139 if (!then_expr) throw std::runtime_error("IfExpr.then_expr is null");1140 if (condition->evaluate(context).to_bool()) {1141 return then_expr->evaluate(context);1142 }1143 if (else_expr) {1144 return else_expr->evaluate(context);1145 }1146 return nullptr;1147 }1148};1149 1150class LiteralExpr : public Expression {1151 Value value;1152public:1153 LiteralExpr(const Location & location, const Value& v)1154 : Expression(location), value(v) {}1155 Value do_evaluate(const std::shared_ptr<Context> &) const override { return value; }1156};1157 1158class ArrayExpr : public Expression {1159 std::vector<std::shared_ptr<Expression>> elements;1160public:1161 ArrayExpr(const Location & location, std::vector<std::shared_ptr<Expression>> && e)1162 : Expression(location), elements(std::move(e)) {}1163 Value do_evaluate(const std::shared_ptr<Context> & context) const override {1164 auto result = Value::array();1165 for (const auto& e : elements) {1166 if (!e) throw std::runtime_error("Array element is null");1167 result.push_back(e->evaluate(context));1168 }1169 return result;1170 }1171};1172 1173class DictExpr : public Expression {1174 std::vector<std::pair<std::shared_ptr<Expression>, std::shared_ptr<Expression>>> elements;1175public:1176 DictExpr(const Location & location, std::vector<std::pair<std::shared_ptr<Expression>, std::shared_ptr<Expression>>> && e)1177 : Expression(location), elements(std::move(e)) {}1178 Value do_evaluate(const std::shared_ptr<Context> & context) const override {1179 auto result = Value::object();1180 for (const auto& [key, value] : elements) {1181 if (!key) throw std::runtime_error("Dict key is null");1182 if (!value) throw std::runtime_error("Dict value is null");1183 result.set(key->evaluate(context), value->evaluate(context));1184 }1185 return result;1186 }1187};1188 1189class SliceExpr : public Expression {1190public:1191 std::shared_ptr<Expression> start, end;1192 SliceExpr(const Location & location, std::shared_ptr<Expression> && s, std::shared_ptr<Expression> && e)1193 : Expression(location), start(std::move(s)), end(std::move(e)) {}1194 Value do_evaluate(const std::shared_ptr<Context> &) const override {1195 throw std::runtime_error("SliceExpr not implemented");1196 }1197};1198 1199class SubscriptExpr : public Expression {1200 std::shared_ptr<Expression> base;