import subprocess
import itertools
import re

def rotl(width, n):
     return "(define-fun rotl_{n} ( (x (_ BitVec {width})) ) (_ BitVec {width})\
             (bvor (bvshl x #x{lshift:x}) (bvlshr x #x{rshift:x})))\n".format(
                     width=width, n=n, lshift=n, rshift=width-n)
     
preface = """
(set-info :smt-lib-version 2.0)
(set-info :status sat)
(set-logic QF_ABV)
"""

postface = """
(check-sat)
;(get-unsat-core)
(exit)
"""

def generate_vars(prefix, width=32, number_of_variables=4):
    variables = ["{prefix:}_{i:}".format(prefix=prefix, i=i) for i in range(number_of_variables)]
    return variables

def generate_vars_smt(variables, width):
    return "\n".join(["(declare-fun {variable} () (_ BitVec {width:}))".format(
        variable=v, width=width) for v in variables 
    ])

def generate_sbox(input_vars, output_vars):
    smt = ""
    return smt

def generate_linear(input_vars, output_vars):
    smt = ""
    return smt


def generate_round(S, round_nr, half_rounds=0):
    """
    S --- SBox ---> N ---- Linear ----> NL
    """
    N = generate_vars("N_{}".format(round_nr), width=4, number_of_variables=320) 
    NL = generate_vars("NL_{}".format(round_nr), width=4, number_of_variables=320)

    smt = ""
    smt +=  generate_sbox(S, N)
    smt += generate_linear(N, NL)

    return [N, NL], smt

def generate_model(rounds=1):
    smt = ""
    plaintext = generate_vars("P_0", 4, 320)
    variables = [plaintext]

    last_state = plaintext
    for r in range(rounds):
        round_vars, round_smt = generate_round_truncated(last_state, r, half_round)
        variables.extend(round_vars)
        smt += round_smt
        last_state = round_vars[-1]
    
    return smt, variables

def generate_conditions(plaintext, ciphertext):
    # First we constrain the output difference to have at least 1 0 or 1 differnce
    smt = """
    """
    return smt

def parse_variables(output):
    regex = re.compile(r"\|([a-zA-Z0-9]+)(?:_([a-zA-Z]{0,1}\d+)){0,1}_(\d+).*#x(\d+)")
    
    matches = regex.findall(output)
    variables = {}

    for m in matches:
        name, r, cell, value = m
        if (r, name) not in variables:
            variables[(r, name)] = [None]*320
        variables[(r, name)][int(cell)] = int(value, 16)
    
    return variables

def print_state(v):
    s = ""
    for i in range(5):
        s += "".join([ f"{x}" if x is not None and x < 2 else "*" for x in v[i : 320 : 5]])
        s += "\n"
    return s

def print_states(variables):
    max_round = max([int(x[0]) for x in variables.keys() if x[0].isdigit()])

    print(variables.keys())
    print(f"Input Diff")
    v = variables[("0", "P")]
    state_string = print_state(v)
    print(f"{state_string}")
    
    for r in range(max_round + 1):
        print(f"{r:2d}")
        v = variables[(f"{r}", "N")]
        state_string = print_state(v)
        print(f"SBox - \n{state_string}")

        v = variables[(f"{r}", "NL")]
        state_string = print_state(v)
        print(f"Linear - \n{state_string}")

    
if __name__ == "__main__":
    rounds = 4
    model_smt, model_vars = generate_model(rounds, 1)

    smt = preface
    smt += generate_vars_smt(itertools.chain(*model_vars), 4)
    smt += model_smt
    smt += generate_conditions(model_vars[0], model_vars[-1])
    smt += postface 
    
    with open("model.smt2", "w") as f:
        f.write(smt)


    p = subprocess.run(["stp", "--print-counterex", "model.smt2"], capture_output=True)
    output = p.stdout.decode("ascii")

    if "sat" not in output or "unsat" in output:
        print(output)
        exit()

    variables = parse_variables(output)
    print_states(variables)
