Team Ai
Apppublic

KBaba7/llama.cpp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
minja.hpp2869 linesDownload Raw Back to common
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;

Showing the first 1,200 of 2869 lines. Download the file for the rest.