Felipe97/llama-cpp-compiled
01.2k
1#pragma once2 3#include "lexer.h"4#include "value.h"5 6#include <cassert>7#include <ctime>8#include <memory>9#include <sstream>10#include <string>11#include <vector>12 13#define JJ_DEBUG(msg, ...) do { if (g_jinja_debug) printf("%s:%-3d : " msg "\n", FILENAME, __LINE__, __VA_ARGS__); } while (0)14 15extern bool g_jinja_debug;16 17namespace jinja {18 19struct statement;20using statement_ptr = std::unique_ptr<statement>;21using statements = std::vector<statement_ptr>;22 23// Helpers for dynamic casting and type checking24template<typename T>25struct extract_pointee_unique {26 using type = T;27};28template<typename U>29struct extract_pointee_unique<std::unique_ptr<U>> {30 using type = U;31};32template<typename T>33bool is_stmt(const statement_ptr & ptr) {34 return dynamic_cast<const T*>(ptr.get()) != nullptr;35}36template<typename T>37T * cast_stmt(statement_ptr & ptr) {38 return dynamic_cast<T*>(ptr.get());39}40template<typename T>41const T * cast_stmt(const statement_ptr & ptr) {42 return dynamic_cast<const T*>(ptr.get());43}44// End Helpers45 46 47// not thread-safe48void enable_debug(bool enable);49 50// for visiting AST nodes51// function signature: void(bool is_leaf, statement * node, pair of <label, children>)52using visitor_pair = std::pair<std::string, std::vector<statement *>>;53using visitor_fn = std::function<void(bool, statement *, std::vector<visitor_pair>)>;54 55struct context {56 std::shared_ptr<std::string> src; // for debugging; use shared_ptr to avoid copying on scope creation57 std::time_t current_time; // for functions that need current time58 59 bool is_get_stats = false; // whether to collect stats60 61 visitor_fn visitor;62 63 // src is optional, used for error reporting64 context(std::string src = "") : src(std::make_shared<std::string>(std::move(src))) {65 env = mk_val<value_object>();66 env->has_builtins = false; // context object has no builtins67 env->insert("true", mk_val<value_bool>(true));68 env->insert("True", mk_val<value_bool>(true));69 env->insert("false", mk_val<value_bool>(false));70 env->insert("False", mk_val<value_bool>(false));71 env->insert("none", mk_val<value_none>());72 env->insert("None", mk_val<value_none>());73 current_time = std::time(nullptr);74 }75 ~context() = default;76 77 context(const context & parent) : context() {78 // inherit variables (for example, when entering a new scope)79 auto & pvar = parent.env->as_ordered_object();80 for (const auto & pair : pvar) {81 set_val(pair.first, pair.second);82 }83 current_time = parent.current_time;84 is_get_stats = parent.is_get_stats;85 src = parent.src;86 }87 88 value get_val(const std::string & name) {89 value default_val = mk_val<value_undefined>(name);90 return env->at(name, default_val);91 }92 93 void set_val(const std::string & name, const value & val) {94 env->insert(name, val);95 }96 97 void set_val(const value & name, const value & val) {98 env->insert(name, val);99 }100 101 void print_vars() const {102 printf("Context Variables:\n%s\n", value_to_json(env, 2).c_str());103 }104 105private:106 value_object env;107};108 109// utils for visiting AST nodes110static std::vector<statement *> stmts_to_ptr(const statements & stmts) {111 std::vector<statement *> children;112 for (const auto & stmt : stmts) {113 children.push_back(stmt.get());114 }115 return children;116}117 118/**119 * Base class for all nodes in the AST.120 */121struct statement {122 size_t pos; // position in source, for debugging123 virtual ~statement() = default;124 virtual std::string type() const { return "Statement"; }125 virtual void visit(context & ctx) { ctx.visitor(true, this, {}); }126 127 // execute_impl must be overridden by derived classes128 virtual value execute_impl(context &) { throw_exec_error(); }129 // execute is the public method to execute a statement with error handling130 value execute(context &);131 132private:133 [[noreturn]] void throw_exec_error() const {134 throw std::runtime_error("cannot exec " + type());135 }136};137 138// Type Checking Utilities139 140template<typename T>141static void chk_type(const statement_ptr & ptr) {142 if (!ptr) return; // Allow null for optional fields143 assert(dynamic_cast<T *>(ptr.get()) != nullptr);144}145 146template<typename T, typename U>147static void chk_type(const statement_ptr & ptr) {148 if (!ptr) return;149 assert(dynamic_cast<T *>(ptr.get()) != nullptr || dynamic_cast<U *>(ptr.get()) != nullptr);150}151 152// Base Types153 154/**155 * Expressions will result in a value at runtime (unlike statements).156 */157struct expression : public statement {158 std::string type() const override { return "Expression"; }159};160 161// Statements162 163struct program : public statement {164 statements body;165 166 program() = default;167 explicit program(statements && body) : body(std::move(body)) {}168 std::string type() const override { return "Program"; }169 [[noreturn]] value execute_impl(context &) override {170 throw std::runtime_error("Cannot execute program directly, use jinja::runtime instead");171 }172};173 174struct if_statement : public statement {175 statement_ptr test;176 statements body;177 statements alternate;178 179 if_statement(statement_ptr && test, statements && body, statements && alternate)180 : test(std::move(test)), body(std::move(body)), alternate(std::move(alternate)) {181 chk_type<expression>(this->test);182 }183 184 std::string type() const override { return "If"; }185 value execute_impl(context & ctx) override;186 void visit(context & ctx) override {187 ctx.visitor(false, this, {188 {"test", {test.get()}},189 {"body", stmts_to_ptr(body)},190 {"alternate", stmts_to_ptr(alternate)}191 });192 }193};194 195struct identifier;196struct tuple_literal;197 198/**199 * Loop over each item in a sequence200 * https://jinja.palletsprojects.com/en/3.0.x/templates/#for201 */202struct for_statement : public statement {203 statement_ptr loopvar; // Identifier | TupleLiteral204 statement_ptr iterable;205 statements body;206 statements default_block; // if no iteration took place207 208 for_statement(statement_ptr && loopvar, statement_ptr && iterable, statements && body, statements && default_block)209 : loopvar(std::move(loopvar)), iterable(std::move(iterable)),210 body(std::move(body)), default_block(std::move(default_block)) {211 chk_type<identifier, tuple_literal>(this->loopvar);212 chk_type<expression>(this->iterable);213 }214 215 std::string type() const override { return "For"; }216 value execute_impl(context & ctx) override;217 void visit(context & ctx) override {218 ctx.visitor(false, this, {219 {"loopvar", {loopvar.get()}},220 {"iterable", {iterable.get()}},221 {"body", stmts_to_ptr(body)},222 {"default_block", stmts_to_ptr(default_block)}223 });224 }225};226 227struct break_statement : public statement {228 std::string type() const override { return "Break"; }229 230 struct signal : public std::exception {231 const char* what() const noexcept override {232 return "Break statement executed";233 }234 };235 236 [[noreturn]] value execute_impl(context &) override {237 throw break_statement::signal();238 }239};240 241struct continue_statement : public statement {242 std::string type() const override { return "Continue"; }243 244 struct signal : public std::exception {245 const char* what() const noexcept override {246 return "Continue statement executed";247 }248 };249 250 [[noreturn]] value execute_impl(context &) override {251 throw continue_statement::signal();252 }253};254 255// do nothing256struct noop_statement : public statement {257 std::string type() const override { return "Noop"; }258 value execute_impl(context &) override {259 return mk_val<value_undefined>();260 }261};262 263struct set_statement : public statement {264 statement_ptr assignee;265 statement_ptr val;266 statements body;267 268 set_statement(statement_ptr && assignee, statement_ptr && value, statements && body)269 : assignee(std::move(assignee)), val(std::move(value)), body(std::move(body)) {270 chk_type<expression>(this->assignee);271 chk_type<expression>(this->val);272 }273 274 std::string type() const override { return "Set"; }275 value execute_impl(context & ctx) override;276 void visit(context & ctx) override {277 ctx.visitor(false, this, {278 {"assignee", {assignee.get()}},279 {"value", {val.get()}},280 {"body", stmts_to_ptr(body)}281 });282 }283};284 285struct macro_statement : public statement {286 statement_ptr name;287 statements args;288 statements body;289 290 macro_statement(statement_ptr && name, statements && args, statements && body)291 : name(std::move(name)), args(std::move(args)), body(std::move(body)) {292 chk_type<identifier>(this->name);293 for (const auto& arg : this->args) chk_type<expression>(arg);294 }295 296 std::string type() const override { return "Macro"; }297 value execute_impl(context & ctx) override;298 void visit(context & ctx) override {299 ctx.visitor(false, this, {300 {"name", {name.get()}},301 {"args", stmts_to_ptr(args)},302 {"body", stmts_to_ptr(body)}303 });304 }305};306 307struct comment_statement : public statement {308 std::string val;309 explicit comment_statement(const std::string & v) : val(v) {}310 std::string type() const override { return "Comment"; }311 value execute_impl(context &) override {312 return mk_val<value_undefined>();313 }314};315 316// Expressions317 318// Represents an omitted expression in a computed member, e.g. `a[]`.319struct blank_expression : public expression {320 std::string type() const override { return "BlankExpression"; }321 value execute_impl(context &) override {322 return mk_val<value_undefined>();323 }324};325 326struct member_expression : public expression {327 statement_ptr object;328 statement_ptr property;329 bool computed; // true if obj[expr] and false if obj.prop330 331 member_expression(statement_ptr && object, statement_ptr && property, bool computed)332 : object(std::move(object)), property(std::move(property)), computed(computed) {333 chk_type<expression>(this->object);334 chk_type<expression>(this->property);335 }336 std::string type() const override { return "MemberExpression"; }337 value execute_impl(context & ctx) override;338 void visit(context & ctx) override {339 ctx.visitor(false, this, {340 {"object", {object.get()}},341 {"property", {property.get()}}342 });343 }344};345 346struct call_expression : public expression {347 statement_ptr callee;348 statements args;349 350 call_expression(statement_ptr && callee, statements && args)351 : callee(std::move(callee)), args(std::move(args)) {352 chk_type<expression>(this->callee);353 for (const auto& arg : this->args) chk_type<expression>(arg);354 }355 std::string type() const override { return "CallExpression"; }356 value execute_impl(context & ctx) override;357 void visit(context & ctx) override {358 ctx.visitor(false, this, {359 {"callee", {callee.get()}},360 {"args", stmts_to_ptr(args)}361 });362 }363};364 365/**366 * Represents a user-defined variable or symbol in the template.367 */368struct identifier : public expression {369 std::string val;370 explicit identifier(const std::string & val) : val(val) {}371 std::string type() const override { return "Identifier"; }372 value execute_impl(context & ctx) override;373};374 375// Literals376 377struct integer_literal : public expression {378 int64_t val;379 explicit integer_literal(int64_t val) : val(val) {}380 std::string type() const override { return "IntegerLiteral"; }381 value execute_impl(context &) override {382 return mk_val<value_int>(val);383 }384};385 386struct float_literal : public expression {387 double val;388 explicit float_literal(double val) : val(val) {}389 std::string type() const override { return "FloatLiteral"; }390 value execute_impl(context &) override {391 return mk_val<value_float>(val);392 }393};394 395struct string_literal : public expression {396 std::string val;397 explicit string_literal(const std::string & val) : val(val) {}398 std::string type() const override { return "StringLiteral"; }399 value execute_impl(context &) override {400 return mk_val<value_string>(val);401 }402};403 404struct array_literal : public expression {405 statements val;406 explicit array_literal(statements && val) : val(std::move(val)) {407 for (const auto& item : this->val) chk_type<expression>(item);408 }409 std::string type() const override { return "ArrayLiteral"; }410 value execute_impl(context & ctx) override {411 auto arr = mk_val<value_array>();412 for (const auto & item_stmt : val) {413 arr->push_back(item_stmt->execute(ctx));414 }415 return arr;416 }417};418 419struct tuple_literal : public expression {420 statements val;421 explicit tuple_literal(statements && val) : val(std::move(val)) {422 for (const auto& item : this->val) chk_type<expression>(item);423 }424 std::string type() const override { return "TupleLiteral"; }425 value execute_impl(context & ctx) override {426 auto arr = mk_val<value_array>();427 for (const auto & item_stmt : val) {428 arr->push_back(item_stmt->execute(ctx));429 }430 return mk_val<value_tuple>(std::move(arr->as_array()));431 }432};433 434struct object_literal : public expression {435 std::vector<std::pair<statement_ptr, statement_ptr>> val;436 explicit object_literal(std::vector<std::pair<statement_ptr, statement_ptr>> && val)437 : val(std::move(val)) {438 for (const auto & pair : this->val) {439 chk_type<expression>(pair.first);440 chk_type<expression>(pair.second);441 }442 }443 std::string type() const override { return "ObjectLiteral"; }444 value execute_impl(context & ctx) override;445};446 447// Complex Expressions448 449/**450 * An operation with two sides, separated by an operator.451 * Note: Either side can be a Complex Expression, with order452 * of operations being determined by the operator.453 */454struct binary_expression : public expression {455 token op;456 statement_ptr left;457 statement_ptr right;458 459 binary_expression(token op, statement_ptr && left, statement_ptr && right)460 : op(std::move(op)), left(std::move(left)), right(std::move(right)) {461 chk_type<expression>(this->left);462 chk_type<expression>(this->right);463 }464 std::string type() const override { return "BinaryExpression"; }465 value execute_impl(context & ctx) override;466 void visit(context & ctx) override {467 ctx.visitor(false, this, {468 {"left", {left.get()}},469 {"right", {right.get()}}470 });471 }472};473 474/**475 * An operation with two sides, separated by the | operator.476 * Operator precedence: https://github.com/pallets/jinja/issues/379#issuecomment-168076202477 */478struct filter_expression : public expression {479 // either an expression or a value is allowed480 statement_ptr operand;481 value_string val; // will be set by filter_statement482 483 statement_ptr filter;484 485 filter_expression(statement_ptr && operand, statement_ptr && filter)486 : operand(std::move(operand)), filter(std::move(filter)) {487 chk_type<expression>(this->operand);488 chk_type<identifier, call_expression>(this->filter);489 }490 491 filter_expression(value_string && val, statement_ptr && filter)492 : val(std::move(val)), filter(std::move(filter)) {493 chk_type<identifier, call_expression>(this->filter);494 }495 496 std::string type() const override { return "FilterExpression"; }497 value execute_impl(context & ctx) override;498 void visit(context & ctx) override {499 ctx.visitor(false, this, {500 {"operand", {operand.get()}},501 {"filter", {filter.get()}}502 });503 }504};505 506struct filter_statement : public statement {507 statement_ptr filter;508 statements body;509 510 filter_statement(statement_ptr && filter, statements && body)511 : filter(std::move(filter)), body(std::move(body)) {512 chk_type<identifier, call_expression>(this->filter);513 }514 std::string type() const override { return "FilterStatement"; }515 value execute_impl(context & ctx) override;516 void visit(context & ctx) override {517 ctx.visitor(false, this, {518 {"filter", {filter.get()}},519 {"body", stmts_to_ptr(body)}520 });521 }522};523 524/**525 * An operation which filters a sequence of objects by applying a test to each object,526 * and only selecting the objects with the test succeeding.527 *528 * It may also be used as a shortcut for a ternary operator.529 */530struct select_expression : public expression {531 statement_ptr lhs;532 statement_ptr test;533 534 select_expression(statement_ptr && lhs, statement_ptr && test)535 : lhs(std::move(lhs)), test(std::move(test)) {536 chk_type<expression>(this->lhs);537 chk_type<expression>(this->test);538 }539 std::string type() const override { return "SelectExpression"; }540 value execute_impl(context & ctx) override {541 auto predicate = test->execute_impl(ctx);542 if (!predicate->as_bool()) {543 return mk_val<value_undefined>();544 }545 return lhs->execute_impl(ctx);546 }547 void visit(context & ctx) override {548 ctx.visitor(false, this, {549 {"lhs", {lhs.get()}},550 {"test", {test.get()}}551 });552 }553};554 555/**556 * An operation with two sides, separated by the "is" operator.557 * NOTE: "value is something" translates to function call "test_is_something(value)"558 */559struct test_expression : public expression {560 statement_ptr operand;561 bool negate;562 statement_ptr test;563 564 test_expression(statement_ptr && operand, bool negate, statement_ptr && test)565 : operand(std::move(operand)), negate(negate), test(std::move(test)) {566 chk_type<expression>(this->operand);567 chk_type<identifier, call_expression>(this->test);568 }569 std::string type() const override { return "TestExpression"; }570 value execute_impl(context & ctx) override;571 void visit(context & ctx) override {572 ctx.visitor(false, this, {573 {"operand", {operand.get()}},574 {"test", {test.get()}}575 });576 }577};578 579/**580 * An operation with one side (operator on the left).581 */582struct unary_expression : public expression {583 token op;584 statement_ptr argument;585 586 unary_expression(token op, statement_ptr && argument)587 : op(std::move(op)), argument(std::move(argument)) {588 chk_type<expression>(this->argument);589 }590 std::string type() const override { return "UnaryExpression"; }591 value execute_impl(context & ctx) override;592 void visit(context & ctx) override {593 ctx.visitor(false, this, {594 {"argument", {argument.get()}}595 });596 }597};598 599struct slice_expression : public expression {600 statement_ptr start_expr;601 statement_ptr stop_expr;602 statement_ptr step_expr;603 604 slice_expression(statement_ptr && start_expr, statement_ptr && stop_expr, statement_ptr && step_expr)605 : start_expr(std::move(start_expr)), stop_expr(std::move(stop_expr)), step_expr(std::move(step_expr)) {606 chk_type<expression>(this->start_expr);607 chk_type<expression>(this->stop_expr);608 chk_type<expression>(this->step_expr);609 }610 std::string type() const override { return "SliceExpression"; }611 [[noreturn]] value execute_impl(context &) override {612 throw std::runtime_error("must be handled by MemberExpression");613 }614 void visit(context & ctx) override {615 ctx.visitor(false, this, {616 {"start_expr", {start_expr.get()}},617 {"stop_expr", {stop_expr.get()}},618 {"step_expr", {step_expr.get()}}619 });620 }621};622 623struct keyword_argument_expression : public expression {624 statement_ptr key;625 statement_ptr val;626 627 keyword_argument_expression(statement_ptr && key, statement_ptr && val)628 : key(std::move(key)), val(std::move(val)) {629 chk_type<identifier>(this->key);630 chk_type<expression>(this->val);631 }632 std::string type() const override { return "KeywordArgumentExpression"; }633 value execute_impl(context & ctx) override;634 void visit(context & ctx) override {635 ctx.visitor(false, this, {636 {"key", {key.get()}},637 {"val", {val.get()}}638 });639 }640};641 642struct spread_expression : public expression {643 statement_ptr argument;644 explicit spread_expression(statement_ptr && argument) : argument(std::move(argument)) {645 chk_type<expression>(this->argument);646 }647 std::string type() const override { return "SpreadExpression"; }648 void visit(context & ctx) override {649 ctx.visitor(false, this, {650 {"argument", {argument.get()}}651 });652 }653};654 655struct call_statement : public statement {656 statement_ptr call;657 statements caller_args;658 statements body;659 660 call_statement(statement_ptr && call, statements && caller_args, statements && body)661 : call(std::move(call)), caller_args(std::move(caller_args)), body(std::move(body)) {662 chk_type<call_expression>(this->call);663 for (const auto & arg : this->caller_args) chk_type<expression>(arg);664 }665 std::string type() const override { return "CallStatement"; }666 value execute_impl(context & ctx) override;667 void visit(context & ctx) override {668 ctx.visitor(false, this, {669 {"call", {call.get()}},670 {"caller_args", stmts_to_ptr(caller_args)},671 {"body", stmts_to_ptr(body)}672 });673 }674};675 676struct ternary_expression : public expression {677 statement_ptr condition;678 statement_ptr true_expr;679 statement_ptr false_expr;680 681 ternary_expression(statement_ptr && condition, statement_ptr && true_expr, statement_ptr && false_expr)682 : condition(std::move(condition)), true_expr(std::move(true_expr)), false_expr(std::move(false_expr)) {683 chk_type<expression>(this->condition);684 chk_type<expression>(this->true_expr);685 chk_type<expression>(this->false_expr);686 }687 std::string type() const override { return "Ternary"; }688 value execute_impl(context & ctx) override {689 value cond_val = condition->execute(ctx);690 if (cond_val->as_bool()) {691 return true_expr->execute(ctx);692 } else {693 return false_expr->execute(ctx);694 }695 }696 void visit(context & ctx) override {697 ctx.visitor(false, this, {698 {"condition", {condition.get()}},699 {"true_expr", {true_expr.get()}},700 {"false_expr", {false_expr.get()}}701 });702 }703};704 705struct raised_exception : public std::exception {706 std::string message;707 raised_exception(const std::string & msg) : message(msg) {}708 const char* what() const noexcept override {709 return message.c_str();710 }711};712 713// Used to rethrow exceptions with modified messages714struct rethrown_exception : public std::exception {715 std::string message;716 rethrown_exception(const std::string & msg) : message(msg) {}717 const char* what() const noexcept override {718 return message.c_str();719 }720};721 722//////////////////////723 724static void gather_string_parts_recursive(const value & val, value_string & parts) {725 // TODO: probably allow print value_none as "None" string? currently this breaks some templates726 if (is_val<value_string>(val)) {727 const auto & str_val = cast_val<value_string>(val)->val_str;728 parts->val_str.append(str_val);729 } else if (is_val<value_int>(val) || is_val<value_float>(val) || is_val<value_bool>(val)) {730 std::string str_val = val->as_string().str();731 parts->val_str.append(str_val);732 } else if (is_val<value_array>(val)) {733 auto items = cast_val<value_array>(val)->as_array();734 for (const auto & item : items) {735 gather_string_parts_recursive(item, parts);736 }737 }738}739 740static std::string render_string_parts(const value_string & parts) {741 std::ostringstream oss;742 for (const auto & part : parts->val_str.parts) {743 oss << part.val;744 }745 return oss.str();746}747 748struct runtime {749 context & ctx;750 explicit runtime(context & ctx) : ctx(ctx) {}751 752 value_array execute(const program & prog) {753 value_array results = mk_val<value_array>();754 for (const auto & stmt : prog.body) {755 value res = stmt->execute(ctx);756 results->push_back(std::move(res));757 }758 return results;759 }760 761 static value_string gather_string_parts(const value & val) {762 value_string parts = mk_val<value_string>();763 gather_string_parts_recursive(val, parts);764 // join consecutive parts with the same type765 auto & p = parts->val_str.parts;766 if (p.empty()) {767 return parts;768 }769 size_t w = 0;770 for (size_t r = 1; r < p.size(); r++) {771 if (p[w].is_input == p[r].is_input) {772 p[w].val += p[r].val;773 } else {774 w++;775 if (w != r) {776 // the guard is needed, self-move leaves the string in an unspecified state777 p[w] = std::move(p[r]);778 }779 }780 }781 p.resize(w + 1);782 return parts;783 }784 785 static std::string debug_dump_program(const program & prog, const std::string & src);786};787 788} // namespace jinja789 