kit

kit
git clone https://git.ryansepassi.com/git/kit.git
Log | Files | Refs | README

pratt_calc_test.c (6462B)


      1 #include "generated_pratt_calc.h"
      2 
      3 #include <ctype.h>
      4 #include <stdio.h>
      5 #include <stdlib.h>
      6 #include <string.h>
      7 
      8 typedef struct Node {
      9     int tk;
     10     long val;
     11 } Node;
     12 
     13 typedef struct {
     14     Node **mem;
     15     size_t nmem, capmem;
     16     int enters, exits, tokens;
     17 } Emb;
     18 
     19 static Node *node_new(Emb *e) {
     20     Node *n = calloc(1, sizeof *n);
     21     if (e->nmem == e->capmem) {
     22         e->capmem = e->capmem ? e->capmem * 2 : 16;
     23         e->mem = realloc(e->mem, e->capmem * sizeof *e->mem);
     24     }
     25     e->mem[e->nmem++] = n;
     26     return n;
     27 }
     28 
     29 static void emb_free(Emb *e) {
     30     for (size_t i = 0; i < e->nmem; i++) free(e->mem[i]);
     31     free(e->mem);
     32 }
     33 
     34 static Node *mk_long(Emb *e, long v) {
     35     Node *n = node_new(e);
     36     n->val = v;
     37     return n;
     38 }
     39 
     40 static long ipow(long base, long exp) {
     41     long out = 1;
     42     while (exp-- > 0) out *= base;
     43     return out;
     44 }
     45 
     46 static long fact(long x) {
     47     long out = 1;
     48     for (long i = 2; i <= x; i++) out *= i;
     49     return out;
     50 }
     51 
     52 static KitGramSem cb_lift(void *ud, KitGramToken t) {
     53     Emb *e = ud;
     54     Node *n = node_new(e);
     55     n->tk = t.kind;
     56     if (t.kind == PRATT_CALC_TOK_NUMBER) {
     57         char b[32];
     58         size_t len = t.len < sizeof b - 1 ? t.len : sizeof b - 1;
     59         memcpy(b, t.lexeme, len);
     60         b[len] = '\0';
     61         n->val = atol(b);
     62     }
     63     return n;
     64 }
     65 
     66 static KitGramSem cb_reduce(void *ud, KitGramRuleId r, int prod, KitGramSem *k, size_t n) {
     67     Emb *e = ud;
     68     (void)n;
     69     switch (r) {
     70     case PRATT_CALC_R_primary:
     71         return prod == 0 ? k[0] : k[1];
     72     case PRATT_CALC_R_expr:
     73         switch (prod) {
     74         case PRATT_CALC_EXPR_PRIMARY:
     75             return k[0];
     76         case PRATT_CALC_EXPR_PREFIX_MINUS:
     77             return mk_long(e, -((Node *)k[1])->val);
     78         case PRATT_CALC_EXPR_PREFIX_PLUS:
     79             return mk_long(e, ((Node *)k[1])->val);
     80         case PRATT_CALC_EXPR_POSTFIX_BANG:
     81             return mk_long(e, fact(((Node *)k[0])->val));
     82         case PRATT_CALC_EXPR_INFIX_PLUS:
     83             return mk_long(e, ((Node *)k[0])->val + ((Node *)k[2])->val);
     84         case PRATT_CALC_EXPR_INFIX_MINUS:
     85             return mk_long(e, ((Node *)k[0])->val - ((Node *)k[2])->val);
     86         case PRATT_CALC_EXPR_INFIX_STAR:
     87             return mk_long(e, ((Node *)k[0])->val * ((Node *)k[2])->val);
     88         case PRATT_CALC_EXPR_INFIX_SLASH:
     89             return mk_long(e, ((Node *)k[0])->val / ((Node *)k[2])->val);
     90         case PRATT_CALC_EXPR_INFIX_CARET:
     91             return mk_long(e, ipow(((Node *)k[0])->val, ((Node *)k[2])->val));
     92         }
     93         break;
     94     }
     95     return NULL;
     96 }
     97 
     98 static void cb_enter(void *ud, KitGramRuleId r, int prod) {
     99     (void)r;
    100     (void)prod;
    101     ((Emb *)ud)->enters++;
    102 }
    103 
    104 static void cb_exit(void *ud, KitGramRuleId r, int prod) {
    105     (void)r;
    106     (void)prod;
    107     ((Emb *)ud)->exits++;
    108 }
    109 
    110 static void cb_token(void *ud, KitGramToken t) {
    111     (void)t;
    112     ((Emb *)ud)->tokens++;
    113 }
    114 
    115 typedef struct { const char *s; size_t i; } Lexer;
    116 
    117 static int lex_next(Lexer *lx, KitGramToken *out) {
    118     while (lx->s[lx->i] && isspace((unsigned char)lx->s[lx->i])) lx->i++;
    119     char c = lx->s[lx->i];
    120     if (!c) return 0;
    121 
    122     out->lexeme = &lx->s[lx->i];
    123     out->len = 1;
    124     out->line = 1;
    125     out->col = (uint32_t)(lx->i + 1);
    126 
    127     if (isdigit((unsigned char)c)) {
    128         size_t j = lx->i;
    129         while (isdigit((unsigned char)lx->s[j])) j++;
    130         out->kind = PRATT_CALC_TOK_NUMBER;
    131         out->len = j - lx->i;
    132         lx->i = j;
    133         return 1;
    134     }
    135 
    136     lx->i++;
    137     switch (c) {
    138     case '+': out->kind = PRATT_CALC_TOK_PLUS;   break;
    139     case '-': out->kind = PRATT_CALC_TOK_MINUS;  break;
    140     case '*': out->kind = PRATT_CALC_TOK_STAR;   break;
    141     case '/': out->kind = PRATT_CALC_TOK_SLASH;  break;
    142     case '^': out->kind = PRATT_CALC_TOK_CARET;  break;
    143     case '!': out->kind = PRATT_CALC_TOK_BANG;   break;
    144     case '(': out->kind = PRATT_CALC_TOK_LPAREN; break;
    145     case ')': out->kind = PRATT_CALC_TOK_RPAREN; break;
    146     default:  out->kind = PRATT_CALC_TOK__COUNT; break;
    147     }
    148     return 1;
    149 }
    150 
    151 static int eval(const char *src, long *result, int *enters, int *exits, int *tokens) {
    152     Emb e;
    153     memset(&e, 0, sizeof e);
    154     KitGramActions acts = {
    155         .reduce = cb_reduce,
    156         .lift_token = cb_lift,
    157         .enter = cb_enter,
    158         .exit = cb_exit,
    159         .on_token = cb_token,
    160     };
    161     KitGramParser ps;
    162     KitGramSlot ctl[256];
    163     KitGramSem vals[256];
    164     KitGramConfig cfg = { .actions = &acts, .ud = &e, .recover = false,
    165                       .ctl_stack = ctl, .ctl_cap = 256,
    166                       .val_stack = vals, .val_cap = 256 };
    167     pratt_calc_parser_init(&ps, &cfg);
    168 
    169     Lexer lx = { src, 0 };
    170     KitGramToken t;
    171     int ok = 1;
    172     while (ok && lex_next(&lx, &t))
    173         if (kit_gram_parser_push(&ps, t) == KIT_GRAM_PARSE_ERROR) ok = 0;
    174     if (ok && kit_gram_parser_finish(&ps) != KIT_GRAM_PARSE_ACCEPT) ok = 0;
    175     if (ok && result) {
    176         Node *r = kit_gram_parser_result(&ps);
    177         *result = r ? r->val : 0;
    178     }
    179     if (enters) *enters = e.enters;
    180     if (exits) *exits = e.exits;
    181     if (tokens) *tokens = e.tokens;
    182     emb_free(&e);
    183     return ok;
    184 }
    185 
    186 static int failures = 0;
    187 
    188 static void check_val(const char *src, long want) {
    189     long got = 0;
    190     int en = 0, ex = 0, tk = 0;
    191     if (!eval(src, &got, &en, &ex, &tk)) {
    192         printf("FAIL  %-18s parse error (wanted %ld)\n", src, want);
    193         failures++;
    194         return;
    195     }
    196     if (got != want) {
    197         printf("FAIL  %-18s = %ld (wanted %ld)\n", src, got, want);
    198         failures++;
    199         return;
    200     }
    201     if (en != ex || en == 0 || tk == 0) {
    202         printf("FAIL  %-18s listener mismatch enter=%d exit=%d tok=%d\n", src, en, ex, tk);
    203         failures++;
    204         return;
    205     }
    206     printf("ok    %-18s = %-5ld [enter=%d exit=%d tok=%d]\n", src, got, en, ex, tk);
    207 }
    208 
    209 static void check_err(const char *src) {
    210     long got = 0;
    211     if (eval(src, &got, NULL, NULL, NULL)) {
    212         printf("FAIL  %-18s parsed = %ld (wanted rejection)\n", src, got);
    213         failures++;
    214         return;
    215     }
    216     printf("ok    %-18s -> rejected\n", src);
    217 }
    218 
    219 int main(void) {
    220     printf("== generated Pratt arithmetic parser ==\n");
    221     check_val("1 + 2 * 3", 7);
    222     check_val("(1 + 2) * 3", 9);
    223     check_val("10 - 2 - 3", 5);
    224     check_val("2 ^ 3 ^ 2", 512);
    225     check_val("-3!", -6);
    226     check_val("2 * -3 + +4", -2);
    227     check_val("3!!", 720);
    228 
    229     check_err("1 +");
    230     check_err("1 2");
    231     check_err("!");
    232 
    233     return failures ? 1 : 0;
    234 }