Team Ai
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes1.2kdownloads
runtime.h789 linesDownload Raw Back to jinja
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