"""Closed JSON and bounded integer AST interpreter for Lab 16; never executes source."""
import ast
import io
import json
import operator
import re
import tokenize

PARAMETERS={'clamp':('x','lo','hi'),'total':('unit','qty','fee')}
TASK=('Repair both functions in toy_math.py using integer expressions.\n'
      'clamp(x, lo, hi): return lo when x < lo, hi when x > hi, otherwise x.\n'
      'Inputs satisfy lo <= hi.\n'
      'total(unit, qty, fee): return the cost of qty units plus one fee.\n'
      'unit, qty, and fee are nonnegative integers.\n'
      'Keep function names and parameters unchanged. Do not change tests.\n')
INITIAL='def clamp(x, lo, hi):\n    return x\n\ndef total(unit, qty, fee):\n    return unit + qty + fee\n'
PUBLIC=(('clamp',(-2,0,5),0),('clamp',(8,0,5),5),('total',(3,4,2),14),('total',(7,0,2),2))
HELDOUT=(('clamp',(-5,-3,-1),-3),('clamp',(0,-3,5),0),('clamp',(5,5,5),5),('clamp',(10,-2,3),3),
         ('total',(0,8,3),3),('total',(6,1,4),10),('total',(2,3,0),6),('total',(9,2,1),19))
BIN={ast.Add:operator.add,ast.Sub:operator.sub,ast.Mult:operator.mul}
COMPARE={ast.Lt:operator.lt,ast.LtE:operator.le,ast.Gt:operator.gt,ast.GtE:operator.ge,ast.Eq:operator.eq,ast.NotEq:operator.ne}

class Denied(ValueError):
    def __init__(self,code): self.code=code;super().__init__(code)


def integer(value):
    if type(value) is not int or abs(value)>1_000_000: raise Denied('integer_limit')
    return value


def expression_text(text):
    if type(text) is not str or not 1<=len(text)<=256 or any(not 32<=ord(c)<=126 for c in text):
        raise Denied('expression_text')
    nesting=0
    for c in text:
        if c=='(': nesting+=1
        elif c==')': nesting-=1
        if not 0<=nesting<=16: raise Denied('expression_nesting')
    if nesting or '#' in text or '\\' in text: raise Denied('expression_text')
    try:
        for token in tokenize.generate_tokens(io.StringIO(text).readline):
            if token.type==tokenize.NUMBER and not re.fullmatch(r'[0-9]{1,7}',token.string):
                raise Denied('integer_token')
            if token.type in (tokenize.STRING,tokenize.COMMENT,tokenize.ERRORTOKEN):
                raise Denied('expression_text')
    except (tokenize.TokenError,IndentationError): raise Denied('expression_text') from None
    return text


def expression(text,function):
    expression_text(text)
    try: root=ast.parse(text.strip(),mode='eval').body
    except (SyntaxError,ValueError,RecursionError): raise Denied('expression_syntax') from None
    stack=[(root,1)]; count=0
    while stack:
        node,depth=stack.pop();count+=1
        if count>128 or depth>16: raise Denied('expression_tree_limit')
        stack.extend((child,depth+1) for child in ast.iter_child_nodes(node))
    def kind(node):
        if type(node) is ast.Constant:
            integer(node.value);return 'integer'
        if type(node) is ast.Name and type(node.ctx) is ast.Load and node.id in PARAMETERS[function]:
            return 'integer'
        if type(node) is ast.BinOp and type(node.op) in BIN:
            if kind(node.left)==kind(node.right)=='integer':return 'integer'
        elif type(node) is ast.UnaryOp and type(node.op) in (ast.UAdd,ast.USub):
            if kind(node.operand)=='integer':return 'integer'
        elif type(node) is ast.Compare and len(node.ops)==len(node.comparators)==1 and type(node.ops[0]) in COMPARE:
            if kind(node.left)==kind(node.comparators[0])=='integer':return 'boolean'
        elif type(node) is ast.IfExp:
            # Validate both branches even when one will never be selected.
            condition,yes,no=kind(node.test),kind(node.body),kind(node.orelse)
            if condition=='boolean' and yes==no=='integer':return 'integer'
        raise Denied('expression_node_or_type')
    if kind(root)!='integer': raise Denied('return_type')
    return root


def source_expressions(source):
    if type(source) is not str or len(source.encode('utf-8'))>4096:raise Denied('source_limit')
    # Source must have only a generated wrapper and two bounded, one-line expressions.
    pattern=r'def clamp\(x, lo, hi\):\n    return ([^\n]+)\n\ndef total\(unit, qty, fee\):\n    return ([^\n]+)\n'
    match=re.fullmatch(pattern,source)
    if not match:raise Denied('source_wrapper')
    values=dict(zip(PARAMETERS,match.groups()))
    for function,text in values.items():expression(text,function)
    try: module=ast.parse(source)
    except SyntaxError:raise Denied('source_wrapper') from None
    if len(module.body)!=2:raise Denied('source_wrapper')
    for node,(name,args) in zip(module.body,PARAMETERS.items()):
        if (type(node) is not ast.FunctionDef or node.name!=name or node.decorator_list
            or node.returns or getattr(node,'type_params',[]) or node.type_comment
            or node.args.posonlyargs or node.args.vararg or node.args.kwarg or node.args.kwonlyargs
            or node.args.defaults or node.args.kw_defaults
            or tuple(a.arg for a in node.args.args)!=args
            or any(a.annotation or a.type_comment for a in node.args.args)
            or len(node.body)!=1 or type(node.body[0]) is not ast.Return):raise Denied('source_wrapper')
    return values


def wrapper(values):
    return f'def clamp(x, lo, hi):\n    return {values["clamp"]}\n\ndef total(unit, qty, fee):\n    return {values["total"]}\n'


def interpret(root,inputs):
    for value in inputs.values():integer(value)
    visits=0
    def walk(node):
        nonlocal visits
        visits+=1
        if visits>256:raise Denied('interpreter_limit')
        if type(node) is ast.Constant:return integer(node.value)
        if type(node) is ast.Name:return integer(inputs[node.id])
        if type(node) is ast.BinOp:
            left,right=integer(walk(node.left)),integer(walk(node.right))
            return integer(BIN[type(node.op)](left,right))
        if type(node) is ast.UnaryOp:
            value=integer(walk(node.operand))
            return integer(value if type(node.op) is ast.UAdd else -value)
        if type(node) is ast.Compare:
            return COMPARE[type(node.ops[0])](integer(walk(node.left)),integer(walk(node.comparators[0])))
        if type(node) is ast.IfExp:
            condition=walk(node.test)
            if type(condition) is not bool:raise Denied('condition_type')
            return integer(walk(node.body if condition else node.orelse))
        raise Denied('expression_node_or_type')
    return integer(walk(root))


def evaluate(source,fixtures,check=lambda:None):
    if len(fixtures)>16:raise Denied('fixture_limit')
    values=source_expressions(source)
    roots={name:expression(text,name) for name,text in values.items()}
    cases=[]
    for name,args,expected in fixtures:
        check(); case={'function':name,'inputs':list(args),'expected':expected}
        try:case['actual']=interpret(roots[name],dict(zip(PARAMETERS[name],args)))
        except Denied as error:case['error']=error.code
        cases.append(case);check()
    return {'passed':sum(c.get('actual')==c['expected'] and 'error' not in c for c in cases),'total':len(cases),'cases':cases}


def parse_proposal(text):
    if type(text) is not str or len(text.encode('utf-8'))>2048:raise Denied('proposal_limit')
    depth=0;quoted=False;escaped=False
    for c in text:
        if quoted:
            if escaped:escaped=False
            elif c=='\\':escaped=True
            elif c=='"':quoted=False
        elif c=='"':quoted=True
        elif c in '{[':
            depth+=1
            if depth>8:raise Denied('json_nesting')
        elif c in '}]':
            depth-=1
            if depth<0:raise Denied('invalid_json')
    def pairs(items):
        result={}
        for key,value in items:
            if key in result:raise Denied('duplicate_key')
            result[key]=value
        return result
    def constant(value):raise Denied('nonfinite_json')
    try:p=json.loads(text,object_pairs_hook=pairs,parse_constant=constant)
    except (json.JSONDecodeError,ValueError,RecursionError) as error:
        if isinstance(error,Denied):raise
        raise Denied('invalid_json') from None
    if type(p) is not dict or type(p.get('action')) is not str:raise Denied('action_contract')
    action=p['action']
    def deny(code):
        error=Denied(code);error.action=action;raise error
    keys={'read':{'action','file'},'replace':{'action','function','old','new','version'},'test':{'action'},'finish':{'action','version'}}
    if action not in keys or set(p)!=keys[action]:deny('action_contract')
    if action=='read' and (type(p['file']) is not str or p['file'] not in ('TASK.txt','toy_math.py')):deny('file_denied')
    if action in ('replace','finish') and (type(p['version']) is not int or not 0<=p['version']<=6):deny('version_type')
    if action=='replace':
        if type(p['function']) is not str or p['function'] not in PARAMETERS:deny('function_denied')
        expression_text(p['old']);expression_text(p['new'])
    return p
