/////////////
// Compile
/////////////

bison -d fold.y
flex fold.l
gcc fold.tab.c lex.yy.c -ll 

/////////////
// Example Inputs
/////////////

1+2+3
3+4+5+6
7
if ((1+2) then (2+3) else (3+4))

/////////////
// fold.y
/////////////

%define parse.error detailed                                                                                                                                                                                                                                                                                                                                                                                                                                                         
 
%{
#include <stdio.h>
#include <stdlib.h>
 
#define NODE_NUMBER 0
#define NODE_ADD    1
#define NODE_SEQ    2
#define NODE_IF     3
#define NODE_FOR    4
#define NODE_ID     5
#define NODE_ASSIGN 6
 
typedef struct AST {
    int type;
    int value;
    char* name;
    struct AST* first;
    struct AST* second;
    struct AST* third;
} AST;
 
static void free_ast(AST* n) {
    if (!n) return; //if n == NULL, return
    free_ast(n->first);
    free_ast(n->second);
    free_ast(n->third);
    free(n);
}
 
static AST* make_node(int type, int value, char* name, AST* first, AST* second, AST* third) {
    AST* n = (AST*)malloc(sizeof(AST));
    n->type = type;
    n->value = value;
    n->name = name;
    n->first = first;
    n->second = second;
    n->third = third;
    return n;
}
static AST* make_seq(AST* l, AST* r){
    if (!r) return l;
    //Dead code elimination
    if (r && (r->type == NODE_NUMBER || r->type == NODE_ID)){
        free_ast(r);
        return l;
    }
    return make_node(NODE_SEQ, 0, NULL, l, r, NULL);
}
 
static AST* make_add(AST* l, AST* r) {
    //Fold away two numbers into one sum
    if (l && r && l->type == NODE_NUMBER && r->type == NODE_NUMBER) {
        int sum = l->value + r->value;
        free_ast(l);
        free_ast(r);
        return make_node(NODE_NUMBER, sum, NULL, NULL, NULL, NULL);
    }
    return make_node(NODE_ADD, 0, NULL, l, r, NULL);
}
 
static AST* make_if(AST* cond, AST* then_branch, AST* else_branch) {
    if (cond && cond->type == NODE_NUMBER) {
        AST* chosen;
        AST* discarded;
        //Fold away always true or always false if statements
        if (cond->value != 0) {
            chosen = then_branch;
            discarded = else_branch;
        } else {
            chosen = else_branch;
            discarded = then_branch;
        }
 
        free_ast(cond);
        free_ast(discarded);
        return chosen;
    }
 
    return make_node(NODE_IF, 0, NULL, cond, then_branch, else_branch);
}
 
// Functional style : (+ a b)
static void pretty(AST* n) {
    if (!n) return;
    if (n->type == NODE_NUMBER) {
        printf("%d", n->value);
    }
    if (n->type == NODE_ID) {
        printf("%s", n->name);
    }
    if (n->type == NODE_ASSIGN) {
        printf("(%s", n->name);
        printf("= ");
        pretty(n->first);
        printf(")");
    }
    if (n->type == NODE_ADD){
        printf("(+ ");
        pretty(n->first);
        printf(" ");
        pretty(n->second);
        printf(")");
    }
 
    if (n->type == NODE_SEQ) {
        pretty(n->first);
        printf("\n");
        pretty(n->second);
    }
    if (n->type == NODE_IF){
        printf("(if ");
        pretty(n->first);
        printf(" ");
        pretty(n->second);
        printf(" ");
        pretty(n->third);
        printf(")");
    }
    if (n->type == NODE_FOR){
        printf("(for ");
        pretty(n->first);
        printf("; ");
        pretty(n->second);
        printf("; ");
        pretty(n->third);
        printf(")");
    }
}
 
static int count_ast(AST* n) {
    if (!n) return 0;
    return 1 + count_ast(n->first) + count_ast(n->second) + count_ast(n->third);
}
 
static void print_ast_ascii(AST* n, const char* prefix, int is_last) {
    char next_prefix[1024];
    int count = 0;
    int i = 0;
 
    if (!n) return;
    printf("%s", prefix);
 
    if (is_last) {
        printf("└── ");
        snprintf(next_prefix, sizeof(next_prefix), "%s    ", prefix);
    } else {
        printf("├── ");
        snprintf(next_prefix, sizeof(next_prefix), "%s│   ", prefix);
    }
 
    if (n->type == NODE_NUMBER) printf("NUMBER %d\n", n->value);
    if (n->type == NODE_ADD)    printf("ADD\n");
    if (n->type == NODE_SEQ)    printf("SEQ\n");
    if (n->type == NODE_IF)     printf("IF\n");
    if (n->type == NODE_FOR)    printf("FOR\n");
    if (n->type == NODE_ID)     printf("ID %s\n", n->name);
    if (n->type == NODE_ASSIGN) printf("ASSIGN %s\n", n->name);
 
    if (n->first)  count++;
    if (n->second) count++;
    if (n->third)  count++;
    if (n->first) { i++; print_ast_ascii(n->first, next_prefix, i == count); }
    if (n->second) { i++; print_ast_ascii(n->second, next_prefix, i == count); }
    if (n->third) { i++; print_ast_ascii(n->third, next_prefix, i == count); }
}
 
int yylex(void);
void yyerror(const char* s);
static AST* root = NULL;
%}
 
%code requires {
    typedef struct AST AST;
}
 
%union {
    int i;
    char* s;
    AST* node;
}
 
%left '+' '='
 
%token FOR IF THEN ELSE EOL
%token <i> NUMBER
%token <s> ID
%type  <node> expr input seq
 
%%
 
input
    : seq         { root = $1; $$ = $1; }
    ;
 
seq
    : seq expr EOL { $$ = make_seq($1,$2); }
    | seq EOL      { $$ = $1; }
    | /* empty */  { $$ = NULL; }
    ;
 
expr
    : NUMBER                              { $$ = make_node(NODE_NUMBER, $1, NULL, NULL, NULL, NULL); }
    | ID                                  { $$ = make_node(NODE_ID, 0, $1, NULL, NULL, NULL); }
    | ID '=' expr                         { $$ = make_node(NODE_ASSIGN, 0, $1, $3, NULL, NULL); }
    | expr '+' expr                       { $$ = make_add($1, $3); }
    | '(' expr ')'                        { $$ = $2; }
    | IF '(' expr THEN expr ELSE expr ')' { $$ = make_if($3, $5, $7); }
    | FOR '(' expr ';' expr ';' expr ')'  { $$ = make_node(NODE_FOR, 0, NULL, $3, $5, $7); }
    ;
 
%%
 
void yyerror(const char* s) {
    fprintf(stderr, "parse error: %s\n", s);
}
 
int main(void) {
    yyparse();
    pretty(root);
    printf("\n\nROOT\n"); print_ast_ascii(root, "", 1);
    printf("\nNodes: %d", count_ast(root));
    free_ast(root);
    printf("\n");
    return 0;
}

////////////
// fold.l
////////////

%{                                                                              
#include <stdlib.h>
#include "fold.tab.h"
%}
 
%%
 
"if"        { return IF; }
"then"      { return THEN; }
"else"      { return ELSE; }
[0-9]+      { yylval.i = atoi(yytext); return NUMBER; }
"+"         { return '+'; }
[ \t\r]+    ;
\n          { return EOL; }
.           { return yytext[0]; }
 
%%
 
int yywrap(void) { return 1; }