Felipe97/llama-cpp-compiled
01.1k
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 