CoolFace
Modelpublic

Felipe97/llama-cpp-compiled

sourceHugging Faceupdated 3d agoView on Hugging Face
0likes1.1kdownloads
parser.cpp603 linesDownload Raw Back to jinja
1#include "lexer.h"2#include "runtime.h"3#include "parser.h"4 5#include <algorithm>6#include <memory>7#include <stdexcept>8#include <string>9#include <vector>10 11#define FILENAME "jinja-parser"12 13namespace jinja {14 15// Helper to check type without asserting (useful for logic)16template<typename T>17static bool is_type(const statement_ptr & ptr) {18    return dynamic_cast<const T*>(ptr.get()) != nullptr;19}20 21class parser {22    const std::vector<token> & tokens;23    size_t current = 0;24 25    std::string source; // for error reporting26 27public:28    parser(const std::vector<token> & t, const std::string & src) : tokens(t), source(src) {}29 30    program parse() {31        statements body;32        while (current < tokens.size()) {33            body.push_back(parse_any());34        }35        return program(std::move(body));36    }37 38    // NOTE: start_pos is the token index, used for error reporting39    template<typename T, typename... Args>40    std::unique_ptr<T> mk_stmt(size_t start_pos, Args&&... args) {41        auto ptr = std::make_unique<T>(std::forward<Args>(args)...);42        assert(start_pos < tokens.size());43        ptr->pos = tokens[start_pos].pos;44        return ptr;45    }46 47private:48    const token & peek(size_t offset = 0) const {49        if (current + offset >= tokens.size()) {50            static const token end_token{token::eof, "", 0};51            return end_token;52        }53        return tokens[current + offset];54    }55 56    const token & next() {57        if (current >= tokens.size()) {58            throw parser_exception("Parser Error: Unexpected EOF", source, tokens.empty() ? 0 : tokens.back().pos);59        }60        return tokens[current++];61    }62 63    token expect(token::type type, const std::string&  error) {64        const auto & t = peek();65        if (t.t != type) {66            throw parser_exception("Parser Error: " + error + " (Got " + t.value + ")", source, t.pos);67        }68        current++;69        return t;70    }71 72    void expect_identifier(const std::string & name) {73        const auto & t = peek();74        if (t.t != token::identifier || t.value != name) {75            throw parser_exception("Expected identifier: " + name, source, t.pos);76        }77        current++;78    }79 80    bool is(token::type type) const {81        return peek().t == type;82    }83 84    bool is_identifier(const std::string & name) const {85        return peek().t == token::identifier && peek().value == name;86    }87 88    bool is_statement(const std::vector<std::string> & names) const {89        if (peek(0).t != token::open_statement || peek(1).t != token::identifier) {90            return false;91        }92        std::string val = peek(1).value;93        return std::find(names.begin(), names.end(), val) != names.end();94    }95 96    statement_ptr parse_any() {97        size_t start_pos = current;98        switch (peek().t) {99            case token::comment:100                return mk_stmt<comment_statement>(start_pos, next().value);101            case token::text:102                return mk_stmt<string_literal>(start_pos, next().value);103            case token::open_statement:104                return parse_jinja_statement();105            case token::open_expression:106                return parse_jinja_expression();107            default:108                throw std::runtime_error("Unexpected token type");109        }110    }111 112    statement_ptr parse_jinja_expression() {113        // Consume {{ }} tokens114        expect(token::open_expression, "Expected {{");115        auto result = parse_expression();116        expect(token::close_expression, "Expected }}");117        return result;118    }119 120    statement_ptr parse_jinja_statement() {121        // Consume {% token122        expect(token::open_statement, "Expected {%");123 124        if (peek().t != token::identifier) {125            throw std::runtime_error("Unknown statement");126        }127 128        size_t start_pos = current;129        std::string name = next().value;130 131        statement_ptr result;132        if (name == "set") {133            result = parse_set_statement(start_pos);134 135        } else if (name == "if") {136            result = parse_if_statement(start_pos);137            // expect {% endif %}138            expect(token::open_statement, "Expected {%");139            expect_identifier("endif");140            expect(token::close_statement, "Expected %}");141 142        } else if (name == "macro") {143            result = parse_macro_statement(start_pos);144            // expect {% endmacro %}145            expect(token::open_statement, "Expected {%");146            expect_identifier("endmacro");147            expect(token::close_statement, "Expected %}");148 149        } else if (name == "for") {150            result = parse_for_statement(start_pos);151            // expect {% endfor %}152            expect(token::open_statement, "Expected {%");153            expect_identifier("endfor");154            expect(token::close_statement, "Expected %}");155 156        } else if (name == "break") {157            expect(token::close_statement, "Expected %}");158            result = mk_stmt<break_statement>(start_pos);159 160        } else if (name == "continue") {161            expect(token::close_statement, "Expected %}");162            result = mk_stmt<continue_statement>(start_pos);163 164        } else if (name == "call") {165            statements caller_args;166            // bool has_caller_args = false;167            if (is(token::open_paren)) {168                // Optional caller arguments, e.g. {% call(user) dump_users(...) %}169                caller_args = parse_args();170                // has_caller_args = true;171            }172            auto callee = parse_primary_expression();173            if (!is_type<identifier>(callee)) throw std::runtime_error("Expected identifier");174 175            auto call_args = parse_args();176            expect(token::close_statement, "Expected %}");177 178            statements body;179            while (!is_statement({"endcall"})) {180                body.push_back(parse_any());181            }182 183            expect(token::open_statement, "Expected {%");184            expect_identifier("endcall");185            expect(token::close_statement, "Expected %}");186 187            auto call_expr = mk_stmt<call_expression>(start_pos, std::move(callee), std::move(call_args));188            result = mk_stmt<call_statement>(start_pos, std::move(call_expr), std::move(caller_args), std::move(body));189 190        } else if (name == "filter") {191            auto filter_node = parse_primary_expression();192            if (is_type<identifier>(filter_node) && is(token::open_paren)) {193                filter_node = parse_call_expression(std::move(filter_node));194            }195            expect(token::close_statement, "Expected %}");196 197            statements body;198            while (!is_statement({"endfilter"})) {199                body.push_back(parse_any());200            }201 202            expect(token::open_statement, "Expected {%");203            expect_identifier("endfilter");204            expect(token::close_statement, "Expected %}");205            result = mk_stmt<filter_statement>(start_pos, std::move(filter_node), std::move(body));206 207        } else if (name == "generation" || name == "endgeneration") {208            // Ignore generation blocks (transformers-specific)209            // See https://github.com/huggingface/transformers/pull/30650 for more information.210            result = mk_stmt<noop_statement>(start_pos);211            ++current;212 213        } else {214            throw std::runtime_error("Unknown statement: " + name);215        }216        return result;217    }218 219    statement_ptr parse_set_statement(size_t start_pos) {220        // NOTE: `set` acts as both declaration statement and assignment expression221        auto left = parse_expression_sequence();222        statement_ptr value = nullptr;223        statements body;224 225        if (is(token::equals)) {226            ++current;227            value = parse_expression_sequence();228        } else {229            // parsing multiline set here230            expect(token::close_statement, "Expected %}");231            while (!is_statement({"endset"})) {232                body.push_back(parse_any());233            }234            expect(token::open_statement, "Expected {%");235            expect_identifier("endset");236        }237        expect(token::close_statement, "Expected %}");238        return mk_stmt<set_statement>(start_pos, std::move(left), std::move(value), std::move(body));239    }240 241    statement_ptr parse_if_statement(size_t start_pos) {242        auto test = parse_expression();243        expect(token::close_statement, "Expected %}");244 245        statements body;246        statements alternate;247 248        // Keep parsing 'if' body until we reach the first {% elif %} or {% else %} or {% endif %}249        while (!is_statement({"elif", "else", "endif"})) {250            body.push_back(parse_any());251        }252 253        if (is_statement({"elif"})) {254            size_t pos0 = current;255            ++current; // consume {%256            ++current; // consume 'elif'257            alternate.push_back(parse_if_statement(pos0)); // nested If258        } else if (is_statement({"else"})) {259            ++current; // consume {%260            ++current; // consume 'else'261            expect(token::close_statement, "Expected %}");262 263            // keep going until we hit {% endif %}264            while (!is_statement({"endif"})) {265                alternate.push_back(parse_any());266            }267        }268        return mk_stmt<if_statement>(start_pos, std::move(test), std::move(body), std::move(alternate));269    }270 271    statement_ptr parse_macro_statement(size_t start_pos) {272        auto name = parse_primary_expression();273        auto args = parse_args();274        expect(token::close_statement, "Expected %}");275        statements body;276        // Keep going until we hit {% endmacro277        while (!is_statement({"endmacro"})) {278            body.push_back(parse_any());279        }280        return mk_stmt<macro_statement>(start_pos, std::move(name), std::move(args), std::move(body));281    }282 283    statement_ptr parse_expression_sequence(bool primary = false) {284        size_t start_pos = current;285        statements exprs;286        exprs.push_back(primary ? parse_primary_expression() : parse_expression());287        bool is_tuple = is(token::comma);288        while (is(token::comma)) {289            ++current; // consume comma290            exprs.push_back(primary ? parse_primary_expression() : parse_expression());291        }292        return is_tuple ? mk_stmt<tuple_literal>(start_pos, std::move(exprs)) : std::move(exprs[0]);293    }294 295    statement_ptr parse_for_statement(size_t start_pos) {296        // e.g., `message` in `for message in messages`297        auto loop_var = parse_expression_sequence(true); // should be an identifier/tuple298        if (!is_identifier("in")) throw std::runtime_error("Expected 'in'");299        ++current; // consume 'in'300 301        // `messages` in `for message in messages`302        auto iterable = parse_expression();303        expect(token::close_statement, "Expected %}");304 305        statements body;306        statements alternate;307 308        // Keep going until we hit {% endfor or {% else309        while (!is_statement({"endfor", "else"})) {310            body.push_back(parse_any());311        }312 313        if (is_statement({"else"})) {314            ++current; // consume {%315            ++current; // consume 'else'316            expect(token::close_statement, "Expected %}");317            while (!is_statement({"endfor"})) {318                alternate.push_back(parse_any());319            }320        }321        return mk_stmt<for_statement>(322            start_pos,323            std::move(loop_var), std::move(iterable),324            std::move(body), std::move(alternate));325    }326 327    statement_ptr parse_expression() {328        // Choose parse function with lowest precedence329        return parse_if_expression();330    }331 332    statement_ptr parse_if_expression() {333        auto a = parse_logical_or_expression();334        if (is_identifier("if")) {335            // Ternary expression336            size_t start_pos = current;337            ++current; // consume 'if'338            auto test = parse_logical_or_expression();339            if (is_identifier("else")) {340                // Ternary expression with else341                size_t pos0 = current;342                ++current; // consume 'else'343                auto false_expr = parse_if_expression(); // recurse to support chained ternaries344                return mk_stmt<ternary_expression>(pos0, std::move(test), std::move(a), std::move(false_expr));345            } else {346                // Select expression on iterable347                return mk_stmt<select_expression>(start_pos, std::move(a), std::move(test));348            }349        }350        return a;351    }352 353    statement_ptr parse_logical_or_expression() {354        auto left = parse_logical_and_expression();355        while (is_identifier("or")) {356            size_t start_pos = current;357            token op = next();358            left = mk_stmt<binary_expression>(start_pos, op, std::move(left), parse_logical_and_expression());359        }360        return left;361    }362 363    statement_ptr parse_logical_and_expression() {364        auto left = parse_logical_negation_expression();365        while (is_identifier("and")) {366            size_t start_pos = current;367            auto op = next();368            left = mk_stmt<binary_expression>(start_pos, op, std::move(left), parse_logical_negation_expression());369        }370        return left;371    }372 373    statement_ptr parse_logical_negation_expression() {374        // Try parse unary operators375        if (is_identifier("not")) {376            size_t start_pos = current;377            auto op = next();378            return mk_stmt<unary_expression>(start_pos, op, parse_logical_negation_expression());379        }380        return parse_comparison_expression();381    }382 383    statement_ptr parse_comparison_expression() {384        // NOTE: membership has same precedence as comparison385        // e.g., ('a' in 'apple' == 'b' in 'banana') evaluates as ('a' in ('apple' == ('b' in 'banana')))386        auto left = parse_additive_expression();387        while (true) {388            token op;389            size_t start_pos = current;390            if (is_identifier("not") && peek(1).t == token::identifier && peek(1).value == "in") {391                op = {token::identifier, "not in", tokens[current].pos};392                ++current; // consume 'not'393                ++current; // consume 'in'394            } else if (is_identifier("in")) {395                op = next();396            } else if (is(token::comparison_binary_operator)) {397                op = next();398            } else break;399            left = mk_stmt<binary_expression>(start_pos, op, std::move(left), parse_additive_expression());400        }401        return left;402    }403 404    statement_ptr parse_additive_expression() {405        auto left = parse_multiplicative_expression();406        while (is(token::additive_binary_operator)) {407            size_t start_pos = current;408            auto op = next();409            left = mk_stmt<binary_expression>(start_pos, op, std::move(left), parse_multiplicative_expression());410        }411        return left;412    }413 414    statement_ptr parse_multiplicative_expression() {415        auto left = parse_test_expression();416        while (is(token::multiplicative_binary_operator)) {417            size_t start_pos = current;418            auto op = next();419            left = mk_stmt<binary_expression>(start_pos, op, std::move(left), parse_test_expression());420        }421        return left;422    }423 424    statement_ptr parse_test_expression() {425        auto operand = parse_filter_expression();426        while (is_identifier("is")) {427            size_t start_pos = current;428            ++current; // consume 'is'429            bool negate = false;430            if (is_identifier("not")) { ++current; negate = true; }431            auto test_id = parse_primary_expression();432            // FIXME: tests can also be expressed like this: if x is eq 3433            if (is(token::open_paren)) test_id = parse_call_expression(std::move(test_id));434            operand = mk_stmt<test_expression>(start_pos, std::move(operand), negate, std::move(test_id));435        }436        return operand;437    }438 439    statement_ptr parse_filter_expression() {440        auto operand = parse_call_member_expression();441        while (is(token::pipe)) {442            size_t start_pos = current;443            ++current; // consume pipe444            auto filter = parse_primary_expression();445            if (is(token::open_paren)) filter = parse_call_expression(std::move(filter));446            operand = mk_stmt<filter_expression>(start_pos, std::move(operand), std::move(filter));447        }448        return operand;449    }450 451    statement_ptr parse_call_member_expression() {452        // Handle member expressions recursively453        auto member = parse_member_expression(parse_primary_expression());454        return is(token::open_paren)455            ? parse_call_expression(std::move(member)) // foo.x()456            : std::move(member);457    }458 459    statement_ptr parse_call_expression(statement_ptr callee) {460        size_t start_pos = current;461        auto expr = mk_stmt<call_expression>(start_pos, std::move(callee), parse_args());462        auto member = parse_member_expression(std::move(expr)); // foo.x().y463        return is(token::open_paren)464            ? parse_call_expression(std::move(member)) // foo.x()()465            : std::move(member);466    }467 468    statements parse_args() {469        // comma-separated arguments list470        expect(token::open_paren, "Expected (");471        statements args;472        while (!is(token::close_paren)) {473            statement_ptr arg;474            // unpacking: *expr475            if (peek().t == token::multiplicative_binary_operator && peek().value == "*") {476                size_t start_pos = current;477                ++current; // consume *478                arg = mk_stmt<spread_expression>(start_pos, parse_expression());479            } else {480                arg = parse_expression();481                if (is(token::equals)) {482                    // keyword argument483                    // e.g., func(x = 5, y = a or b)484                    size_t start_pos = current;485                    ++current; // consume equals486                    arg = mk_stmt<keyword_argument_expression>(start_pos, std::move(arg), parse_expression());487                }488            }489            args.push_back(std::move(arg));490            if (is(token::comma)) {491                ++current; // consume comma492            }493        }494        expect(token::close_paren, "Expected )");495        return args;496    }497 498    statement_ptr parse_member_expression(statement_ptr object) {499        size_t start_pos = current;500        while (is(token::dot) || is(token::open_square_bracket)) {501            auto op = next();502            bool computed = op.t == token::open_square_bracket;503            statement_ptr prop;504            if (computed) {505                prop = parse_member_expression_arguments();506                expect(token::close_square_bracket, "Expected ]");507            } else {508                prop = parse_primary_expression();509            }510            object = mk_stmt<member_expression>(start_pos, std::move(object), std::move(prop), computed);511        }512        return object;513    }514 515    statement_ptr parse_member_expression_arguments() {516        // NOTE: This also handles slice expressions colon-separated arguments list517        // e.g., ['test'], [0], [:2], [1:], [1:2], [1:2:3]518        statements slices;519        bool is_slice = false;520        size_t start_pos = current;521        while (!is(token::close_square_bracket)) {522            if (is(token::colon)) {523                // A case where a default is used524                // e.g., [:2] will be parsed as [undefined, 2]525                slices.push_back(nullptr);526                ++current; // consume colon527                is_slice = true;528            } else {529                slices.push_back(parse_expression());530                if (is(token::colon)) {531                    ++current; // consume colon after expression, if it exists532                    is_slice = true;533                }534            }535        }536        if (is_slice) {537            statement_ptr start = slices.size() > 0 ? std::move(slices[0]) : nullptr;538            statement_ptr stop = slices.size() > 1 ? std::move(slices[1]) : nullptr;539            statement_ptr step = slices.size() > 2 ? std::move(slices[2]) : nullptr;540            return mk_stmt<slice_expression>(start_pos, std::move(start), std::move(stop), std::move(step));541        }542        if (slices.empty()) {543            return mk_stmt<blank_expression>(start_pos);544        }545        return std::move(slices[0]);546    }547 548    statement_ptr parse_primary_expression() {549        size_t start_pos = current;550        auto t = next();551        switch (t.t) {552            case token::numeric_literal:553                if (t.value.find('.') != std::string::npos) {554                    return mk_stmt<float_literal>(start_pos, std::stod(t.value));555                } else {556                    return mk_stmt<integer_literal>(start_pos, std::stoll(t.value));557                }558            case token::string_literal: {559                std::string val = t.value;560                while (is(token::string_literal)) {561                    val += next().value;562                }563                return mk_stmt<string_literal>(start_pos, val);564            }565            case token::identifier:566                return mk_stmt<identifier>(start_pos, t.value);567            case token::open_paren: {568                auto expr = parse_expression_sequence();569                expect(token::close_paren, "Expected )");570                return expr;571            }572            case token::open_square_bracket: {573                statements vals;574                while (!is(token::close_square_bracket)) {575                    vals.push_back(parse_expression());576                    if (is(token::comma)) ++current;577                }578                ++current;579                return mk_stmt<array_literal>(start_pos, std::move(vals));580            }581            case token::open_curly_bracket: {582                std::vector<std::pair<statement_ptr, statement_ptr>> pairs;583                while (!is(token::close_curly_bracket)) {584                    auto key = parse_expression();585                    expect(token::colon, "Expected :");586                    pairs.push_back({std::move(key), parse_expression()});587                    if (is(token::comma)) ++current;588                }589                ++current;590                return mk_stmt<object_literal>(start_pos, std::move(pairs));591            }592            default:593                throw std::runtime_error("Unexpected token: " + t.value + " of type " + std::to_string(t.t));594        }595    }596};597 598program parse_from_tokens(const lexer_result & lexer_res) {599    return parser(lexer_res.tokens, lexer_res.source).parse();600}601 602} // namespace jinja603