import math
import sys
import os
import re
from LambdaParser import parser

counter = 0

def free_variables(tree):
  if tree[0] == "name":
    return {tree[1]}
  elif tree[0] == "num":
    return set()
  elif tree[0] == "lambda":
    t = free_variables(tree[2])
    t.discard(tree[1])
    return t
  else:  # must be  "op" or "apply"
    return free_variables(tree[1]).union(free_variables(tree[2]))

def alpha_replace(tree,oldvar,newvar):
  if tree[0] == "name":
    if tree[1] == oldvar:
      return ["name",newvar]
    else:
      return ["name",tree[1]]
  elif tree[0] == "num":
    return ["num",tree[1]]
  elif tree[0] == "lambda":
    if tree[1] == oldvar:
      return ["lambda",oldvar,tree[2]]
    else:
      return ["lambda",tree[1],alpha_replace(tree[2],oldvar,newvar)]
  else:  # must be  "op" or "apply"
    return [tree[0],alpha_replace(tree[1],oldvar,newvar),alpha_replace(tree[2],oldvar,newvar)]

def alpha_convert(tree,var):
  if tree[0] == 'lambda':
    return ['lambda',var,alpha_replace(tree[2],tree[1],var)]
  else:
    return tree

def substitute(tree,var,val):
  if tree[0] == "name":
    if tree[1] == var:
      return val
    else:
      return tree
  elif tree[0] == "num":
    return tree
  elif tree[0] == "lambda":
    if tree[1] == var:
      return tree
    elif tree[1] not in free_variables(val):
      return ['lambda',tree[1],substitute(tree[2],var,val)]
    else:
      global counter
      newvar = '_'+str(counter)
      counter += 1
      [a,b,new_body] = alpha_convert(tree,newvar)
      return ['lambda',newvar,substitute(new_body,var,val)]
  else:  # must be  "op" or "apply"
    return [tree[0],substitute(tree[1],var,val),substitute(tree[2],var,val)]

#
# Finds one beta-reduction; performs it and returns (True,newTree)
# If no beta-reduction found returns (False,tree)
# Uses normal order, i.e. left-most beta-reduction
#
def beta_reduction(tree):
  if tree[0] == "name" or tree[0] == "num":
    return (False,tree) 
  elif tree[0] == "lambda":
    (b,t) = beta_reduction(tree[2])
    if b:
      return (True,["lambda",tree[1],t])
    else:
      return (False,tree)    
  elif tree[0] == "apply":
    if tree[1][0] == "lambda":
      return (True,substitute(tree[1][2],tree[1][1],tree[2]))
    else:
      (b,t) = beta_reduction(tree[1])
      if b:
        return (True,["apply",t,tree[2]])
      else:
        (b,t) = beta_reduction(tree[2])
        if b:
          return (True,["apply",tree[1],t])
        else:
          return (False,tree)
  else: # must be "op"
    (b,t) = beta_reduction(tree[1])
    if b:
      return (True,[tree[0],t,tree[2]])
    else:
      (b,t) = beta_reduction(tree[2])
      if b:
        return (True,[tree[0],tree[1],t])
      else:
        return (False,tree)

def numeric_evaluate(tree):
  if tree[0] == "name" or tree[0] == "num":
    return tree
  elif tree[0] == "lambda":
    return ["lambda",tree[1],numeric_evaluate(tree[2])]
  elif tree[0] == "apply":
    return ["apply",numeric_evaluate(tree[1]),numeric_evaluate(tree[2])]
  else: # must be "op"
    val1 = numeric_evaluate(tree[1])
    val2 = numeric_evaluate(tree[2])
    if val1[0] == "num" and val2[0] == "num":
      if tree[0] == "+":
        return ["num",val1[1]+val2[1]]
      elif tree[0] == "-":
        return ["num",val1[1]-val2[1]]
      elif tree[0] == "*":
        return ["num",val1[1]*val2[1]]
      elif val2[1] != 0:
        return ["num",val1[1]/val2[1]]
      else:
        return ["num",math.nan]
    elif val1[0] == "num" and val2[0] != "num":
      return [tree[0],["num",val1[1]],val2]
    elif val1[0] != "num" and val2[0] == "num":
      return [tree[0],val1,["num",val2[1]]]
    else:
      return [tree[0],val1,val2]

def evaluate(tree,debug):
  while True:
    (b,tree) = beta_reduction(tree)
    if not b:
      break
    #Turned off single step
    if debug:
      print("BETA: ",to_string(tree))
  return tree

def to_string(tree):
  s = ""
  if tree[0] == "name":
    s += tree[1]
  elif tree[0] == "num":
    s += str(tree[1])
  elif tree[0] == "lambda": 
    s += "(LAMBDA "+tree[1]+" "+to_string(tree[2])+")"
  elif tree[0] == "apply":
    s += "("+to_string(tree[1])+" "+to_string(tree[2])+")"
  else: # must be  "op"
    s += "("+tree[0]+" "+to_string(tree[1])+" "+to_string(tree[2])+")"
  return s

def read_input():
  result = ''
  while True:
    data = input('LAMBDA> ').strip()
    if ';' in data:
      i = data.index(';')
      result += data[0:i + 1]
      break
    else:
      result += data + ' '
  return result

def tokenize_string(input_str):
    # This pattern matches:
    # 1. '(' or ')'      -> Parentheses
    # 2. '\s+'           -> One or more whitespace characters
    # 3. '[^()\s]+'      -> Sequences of characters that are NOT parens or whitespace (atoms/words)
    # 4. ';'             -> Semicolons (often used as terminators in Lisp-like languages)
    pattern = r"([()]|\s+|[^()\s;]+|;)"

    # re.split includes the capturing groups in the result list.
    # We filter out None or empty strings that might result from the split.
    tokens = [t for t in re.split(pattern, input_str) if t]

    return tokens

def replace_names(expr,d):
  while True:
    changed = False
    tokens = tokenize_string(expr)
    result = []
    for token in tokens:
      if token in d:
        result.append(d[token])
        changed = True
      else:
        result.append(token)
    if changed:
      expr = "".join(result)
    else:
      break
  return "".join(result)
      
def circular_definitions(names):
    return False

def main():
  debug = (len(sys.argv) == 2 and sys.argv[1] == "-d")
  names = {}
  file_path = os.path.expanduser('~/.startup-lambda')
  if os.path.isfile(file_path):
    with open(file_path) as f:
      lines = f.readlines()
    for line in lines:
      if line.strip() == "":
        continue
      ss = line.split("=")
      name = ss[0].strip()
      expr = ss[1].strip()
      if name in names.keys():
        print("Invalid name:",name)
        continue
      names[name] = expr
  #print(names)
  if circular_definitions(names):
      print("There are circular definitions in ~/.startup-lambda")
      sys.exit(1)
  while True:
    # Read input
    data = read_input()
    if data == 'exit;':
      break
    # We have an input
    data = replace_names(data,names)
    #print(data)
    # Parse
    try:
      tree = parser.parse(data.upper())
    except Exception as inst:
      print(inst.args[0])
      continue
    # We have a successful parse
    # Evaluate
    try:
      if tree[0] == 'subst':
        print(to_string(tree[1]))
        t = substitute(tree[1],tree[2],tree[3])
        print(to_string(t))
        #print("Substitute feature not implemented")
      elif tree[0] == 'freevars':
        print("Free variables: ", free_variables(tree[1]))
        #print("Free variables feature not implemented")
      elif tree[0] == 'alpha':
        if tree[1][0] == 'lambda':
          newtree = alpha_convert(tree[1],tree[2])
          print(to_string(newtree))
        else:
          print("Cannot apply alpha-conversion; top node must be a LAMBDA")
        #print("Alpha equivalence not implemented")
      else:
        print("\nINPUT: ",to_string(tree))
        answer = evaluate(tree,debug)
        print("BETA REDUCT: ",to_string(answer))
        print("\nANSWER: ",to_string(numeric_evaluate(answer)),"\n")
    except Exception as inst:
      print(inst.args[0])
      continue

main()
