#!/usr/bin/env python3

import sys
import pathlib

bcode_header_fmt = '''
#ifndef BCGEN_BCODE_H
#define BCGEN_BCODE_H

#include <stdio.h>
#include <stddef.h>

typedef {self.ubcval_type} ubcval_t;
typedef {self.sbcval_type} sbcval_t;
typedef float fbcval_t;
typedef double dbcval_t;
typedef size_t bcreloc_t;

typedef union {{
	ubcval_t r;
	sbcval_t s;
	fbcval_t f;
	dbcval_t d;
}} bcval_t;

struct run_state {{
	void* (*op);
        bcval_t *imm[{self.i}];
}};

struct compile_state {{
	struct run_state s;
	void* (*labels);
	size_t size;
	size_t pc;
}};

{self.select_header}

void run(struct compile_state *cs, ubcval_t ri[{self.r}], dbcval_t fi[{self.f}]);
void init(struct compile_state *cs);
void end(struct compile_state *cs);
void destroy(struct compile_state *cs);

bcval_t label(struct compile_state *cs);
void patch(struct compile_state *cs, bcreloc_t reloc, size_t i, bcval_t imm);

#endif /* BCGEN_BCODE_H */
'''

bcode_src_fmt = '''
#include <stdlib.h>
#include <assert.h>
#include <stdint.h>
#include <stddef.h>
#include <stdbool.h>

#include "{self.name}.h"

{self.head}

enum bcode_insn {{
	END,
    {self.enum}
}};

static void append(struct compile_state *cs, enum bcode_insn label,
                    size_t n, bcval_t imm[])
{{
	assert(cs->labels);

	if (cs->pc == cs->size) {{
		cs->size *= 2;
		/* better memory checks could be a good idea */
		cs->s.op = realloc(cs->s.op, cs->size * sizeof(*cs->s.op));
		assert(cs->s.op);

                for (size_t i = 0; i < {self.i}; ++i) {{
                    cs->s.imm[i] = realloc(cs->s.imm[i],
                        cs->size * sizeof(*cs->s.imm[i]));
                    assert(cs->s.imm[i]);
                }}
	}}

	struct run_state *s = &cs->s;
	s->op[cs->pc] = cs->labels[label];
        for (size_t i = 0; i < n; ++i)
	    s->imm[i][cs->pc] = imm[i];

	cs->pc++;
}}

#define RELOC cs->pc

{self.select}
{self.macro}

#define IMM0  s.imm[0][pc].r
#define SIMM0 s.imm[0][pc].s
#define FIMM0 s.imm[0][pc].f
#define DIMM0 s.imm[0][pc].d

#define IMM(i)  s.imm[(i)][pc].r
#define SIMM(i) s.imm[(i)][pc].s
#define FIMM(i) s.imm[(i)][pc].f
#define DIMM(i) s.imm[(i)][pc].d

#define NEXT_INSN goto *s.op[++pc]
#define JUMP(i) goto *s.op[pc = (i)]

static void _run(struct compile_state *cs, bool init,
    ubcval_t ri[{self.r}],
    dbcval_t fi[{self.f}])
{{
        // silence warnings if zero registers defined
        (void)ri;
        (void)fi;
	static void *labels[] = {{
		[END] = &&end,
        {self.static}
	}};

	if (init) {{
		cs->labels = labels;
		return;
	}}

{self.regs}
{self.extra}

	const struct run_state s = cs->s;
	size_t pc = 0;
	JUMP(pc);

{self.body}

end:
        {self.rregs}
	return;
}}

void run(struct compile_state *cs, ubcval_t ri[{self.r}], dbcval_t fi[{self.f}])
{{
	_run(cs, false, ri, fi);
}}

void init(struct compile_state *cs)
{{
	cs->labels = NULL;
	cs->size = 16;

        for (size_t i = 0; i < {self.i}; ++i) {{
            cs->s.imm[i] = malloc(cs->size * sizeof(*cs->s.imm));
            assert(cs->s.imm[i]);
        }}

	cs->s.op = malloc(cs->size * sizeof(*cs->s.op));
	assert(cs->s.op);

	cs->pc = 0;
	_run(cs, true, NULL, NULL);
}}

void end(struct compile_state *cs)
{{
	append(cs, END, 0, NULL);
}}

bcval_t label(struct compile_state *cs)
{{
	return (bcval_t){{.r = cs->pc}};
}}

void patch(struct compile_state *cs, bcreloc_t reloc, size_t i, bcval_t imm)
{{
	cs->s.imm[i][reloc] = imm;
}}

void destroy(struct compile_state *cs)
{{
	free(cs->s.op);

        for (size_t i = 0; i < {self.i}; ++i)
	    free(cs->s.imm[i]);
}}
'''

def check_fmt(fmt, imms):
    # todo: should sign value be encoded in the registers?
    assert(len(fmt) == 5 + imms)
    assert(fmt[0] == '_' or fmt[0] == 'r')

    for i in range(1, 5):
        c = fmt[i]
        assert(c == 'R' or c == 'F' or c == '_')

    for i in range(5, 5 + imms):
        c = fmt[i]
        assert(c == '_' or c == 'I' or c == 'D' or c == 'S')

def param(fmt, i):
    c = fmt[i + 1]
    return c + str(i)

def gen_label(name, r0, r1, r2, r3):
    return '{}_{}{}{}{}'.format(name, r0, r1, r2, r3)

def gen_enum(name, r0, r1, r2, r3):
    return gen_label(name, r0, r1, r2, r3).upper()

def num_imms(fmt):
    n = 0
    for p in fmt[5:]:
        if p == 'I' or p == 'S' or p == 'D':
            n += 1

    return n

def gen_hash(fmt):
    n = 0
    i = 0
    r = ""
    for p in fmt[1:]:
        if p == 'R' or p == 'F':
            r += ' + (a' + str(n) + ' << (' + str(i) + ' * 8))'
            n += 1
            i += 1

        elif p == '_':
            i += 1
            continue

        else:
            break

    return r


def gen_params(fmt):
    n = 0
    r = ""
    for p in fmt[1:]:
        if p == 'R' or p == 'F':
            # todo: could be useful to have different types for general purpose
            # registers and floating point registers, might avoid bugs that
            # would otherwise be pretty annoying to track down.
            r += ', size_t a' + str(n)
            n += 1

        elif p == '_':
            continue

        else:
            break

    n = 0
    i = ""
    for p in fmt[5:]:
        if p == 'I':
            i += ', ubcval_t i' + str(n)
            n += 1
        elif p == 'S':
            i = ', sbcval_t i' + str(n)
            n += 1
        elif p == 'D':
            i = ', fbcval_t i' + str(n)
            n += 1

    return '(struct compile_state *cs' + r + i + ')'

def gen_signature(name, fmt):
    ret = 'void '
    if fmt[0] == 'r':
        ret = 'bcreloc_t '

    n = 'select_' + name
    p = gen_params(fmt)
    return ret + n + p

def gen_cond(r, i):
    if r == '':
        return ' + 0'

    return ' + (' + r[1:] + ' << (' + str(i) + ' * 8))'

class BCGen:
    def __init__(self, r = 3, f = 3, i = 1,
            ubcval_type = 'unsigned long',
            sbcval_type = 'signed long',
            extra='', head=''):

        self.r = r
        self.f = f
        self.i = i

        self.ubcval_type = ubcval_type
        self.sbcval_type = sbcval_type

        self.extra = extra
        self.head = head

        self.regs = []
        self.body = []
        self.enum = []
        self.rregs = []
        self.macro = []
        self.static = []
        self.select = []
        self.select_header = []

    def gen_exec_instance(self, name, r0, r1, r2, r3):
        label = gen_label(name, r0, r1, r2, r3) + ': '
        macro = '{}({}, {}, {}, {}); '.format(name.upper(), r0, r1, r2, r3)
        next = 'NEXT_INSN;'
        self.body.append(label + macro + next + '\n')

    # gen should maybe just generate the string and print_ would print it?
    def gen_static_instance(self, name, r0, r1, r2, r3):
        label = gen_label(name, r0, r1, r2, r3)
        enum = gen_enum(name, r0, r1, r2, r3)
        self.static.append('[' + enum + '] = &&' + label + ',\n')

    def gen_select_instance(self, name, r0, r1, r2, r3):
        c0 = gen_cond(r0, 0)
        c1 = gen_cond(r1, 1)
        c2 = gen_cond(r2, 2)
        c3 = gen_cond(r3, 3)
        cond = c0 + c1 + c2 + c3
        action = 'label = ' + gen_enum(name, r0, r1, r2, r3) + ';'
        self.select.append('case ' + cond + ': {' + action + '}; break;\n')

    def gen_enum_instance(self, name, r0, r1, r2, r3):
        enum = gen_enum(name, r0, r1, r2, r3)
        self.enum.append(enum + ',\n')

    def gen_macro(self, name, fmt, body):
        name = name.upper()

        p0 = param(fmt, 0)
        p1 = param(fmt, 1)
        p2 = param(fmt, 2)
        p3 = param(fmt, 3)

        macro = '#define {}({}, {}, {}, {}) {}\n'.format(
                name, p0, p1, p2, p3, body.replace('\n', ''))

        self.macro.append(macro)

    def place(self, fmt, i):
        c = fmt[i + 1]
        if c == 'R':
            return ['r' + str(i) for i in range(self.r)]

        if c == 'F':
            return ['f' + str(i) for i in range(self.f)]

        return [""]

    def gen_body(self, name, fmt):
        for r0 in self.place(fmt, 0):
            for r1 in self.place(fmt, 1):
                for r2 in self.place(fmt, 2):
                    for r3 in self.place(fmt, 3):
                        self.gen_exec_instance(name, r0, r1, r2, r3)

    def gen_static(self, name, fmt):
        for r0 in self.place(fmt, 0):
            for r1 in self.place(fmt, 1):
                for r2 in self.place(fmt, 2):
                    for r3 in self.place(fmt, 3):
                        self.gen_static_instance(name, r0, r1, r2, r3)

    def gen_select(self, name, fmt):
        sign = gen_signature(name, fmt)
        self.select_header.append(sign + ';\n')

        self.select.append(sign + '\n')
        self.select.append('{\nenum bcode_insn label = 0;\n')
        if fmt[0] == 'r':
            self.select.append('bcreloc_t reloc = RELOC;\n')

        calc_hash = gen_hash(fmt);
        self.select.append('unsigned long hash = ' + calc_hash + ';\n');
        self.select.append('switch (hash) {\n')

        # we could potentially just calculate the wanted label directly and skip
        # the huge number of if statements, potentially giving us a slight
        # compile time speedup
        for r0 in self.place(fmt, 0):
            for r1 in self.place(fmt, 1):
                for r2 in self.place(fmt, 2):
                    for r3 in self.place(fmt, 3):
                        self.gen_select_instance(name, r0, r1, r2, r3)

        self.select.append('default: fprintf(stderr,' +
                           '"invalid args to ' + name + '"); abort(); }\n')

        n = num_imms(fmt)
        imms = "NULL"
        if n > 0:
            s = '{i' + str(0) + '}'
            for i in range(n - 1):
                s += ', {i' + str(i) + '}'

            self.select.append('bcval_t imms[{n}] = {{{s}}};\n'.format(n=n, s=s))
            imms = "imms"


        self.select.append('append(cs, label, {n}, {imms});\n'.format(n=n, imms=imms))

        if fmt[0] == 'r':
            self.select.append('return reloc;\n')

        self.select.append('}\n')

    def gen_enum(self, name, fmt):
        for r0 in self.place(fmt, 0):
            for r1 in self.place(fmt, 1):
                for r2 in self.place(fmt, 2):
                    for r3 in self.place(fmt, 3):
                        self.gen_enum_instance(name, r0, r1, r2, r3)


    def rule(self, name, fmt, body):
        check_fmt(fmt, self.i)
        self.gen_macro(name, fmt, body)
        self.gen_static(name, fmt)
        self.gen_select(name, fmt)
        self.gen_body(name, fmt)
        self.gen_enum(name, fmt)

    def write(self, name='bcode'):
        gregs = ['ubcval_t r' + str(i) + ' = ri[' + str(i) + '];\n'
                    for i in range(self.r)]
        fregs = ['dbcval_t f' + str(i) + ' = fi[' + str(i) + '];\n'
                    for i in range(self.f)]

        rgregs = ['ri[' + str(i) + '] = r' + str(i) + ';\n'
                    for i in range(self.r)]
        rfregs = ['fi[' + str(i) + '] = f' + str(i) +';\n'
                    for i in range(self.f)]

        for g in gregs:
            self.regs += g

        for f in fregs:
            self.regs += f

        for g in rgregs:
            self.rregs += g

        for f in rfregs:
            self.rregs += f

        with open(name + '.h', 'w') as f:
            self.select_header = ''.join(self.select_header)
            f.write(bcode_header_fmt.format(self=self))

        with open(name + '.c', 'w') as f:
            # join arrays into single strings
            self.name = name
            self.head = ''.join(self.head)
            self.enum = ''.join(self.enum)
            self.select = ''.join(self.select)
            self.macro = ''.join(self.macro)
            self.static = ''.join(self.static)
            self.regs = ''.join(self.regs)
            self.extra = ''.join(self.extra)
            self.body = ''.join(self.body)
            self.rregs = ''.join(self.rregs)

            f.write(bcode_src_fmt.format(self=self))

assert(len(sys.argv) == 2)
exec(open(sys.argv[1]).read())
