from multiprocessing import Pool
import math, numpy

#------ Arithmetic Functions ------#
def isSquare(apositiveint): # Returns True if apositiveint has an integer sqrt
  x = apositiveint // 2
  seen = set([x])
  while x * x != apositiveint:
    x = (x + (apositiveint // x)) // 2
    if x in seen: return False
    seen.add(x)
  return True

def withinTwoLimits(functionValue, lowerLimit, upperLimit):
   if functionValue < lowerLimit:
      return lowerLimit
   if functionValue > upperLimit:
      return upperLimit
   return functionValue

def aboveLimit(functionValue, lowerLimit):
    if functionValue < lowerLimit:
      return lowerLimit
    return functionValue

def belowLimit(functionValue, upperLimit):
   if functionValue > upperLimit:
      return upperLimit
   return functionValue

def BayesianPosterior(PH, PEH, PnotH, PEnotH):
   return (PH*PEH)/(PH*PEH+PnotH*PEnotH)

def QuadraticRoot(a, b, c):
   discriminant = math.sqrt(b**2-4*a*c)

   return ((-b-discriminant)/2*a, (-b+discriminant)/2*a)


#------ Comms Functions ------#
def GetProbabilities(signals, priors):
   # List elements: acceleration, eye_contact, gesture
   AtCo = sorted([1]*(1+len(signals.keys())))
   AtPu = sorted([1]*(1+len(signals.keys())))
   DiCo = sorted([1]*(1+len(signals.keys())))
   DiPu = sorted([1]*(1+len(signals.keys())))

   i=0

   #Priors
   AtCo[i] = priors.get("Attentive")*priors.get("Cooperative")
   AtPu[i] = priors.get("Attentive")*priors.get("Punitive")
   DiCo[i] = priors.get("Distracted")*priors.get("Cooperative")
   DiPu[i] = priors.get("Distracted")*priors.get("Punitive")
   i+=1

   if 'acceleration' in signals.keys():
      # Acceleration - observed from "real" data.
      AtCo[i] = 0.55 if signals.get('acceleration') == 1 else (0.05 if signals.get('acceleration') == 0 else 0.40)
      AtPu[i] = 0.30 if signals.get('acceleration') == 1 else (0.05 if signals.get('acceleration') == 0 else 0.65)
      DiCo[i] = 0.25 if signals.get('acceleration') == 1 else (0.55 if signals.get('acceleration') == 0 else 0.20)
      DiPu[i] = 0.15 if signals.get('acceleration') == 1 else (0.55 if signals.get('acceleration') == 0 else 0.30)
      i+=1

   if 'eyes' in signals.keys():
      # Eye Contact
      AtCo[i] = 0.90 if signals.get('eyes') == 1 else 0.10
      AtPu[i] = 0.90 if signals.get('eyes') == 1 else 0.10
      DiCo[i] = 0.05 if signals.get('eyes') == 1 else 0.95
      DiPu[i] = 0.05 if signals.get('eyes') == 1 else 0.95
      i+=1

   if 'signal' in signals.keys():
      # Explicit Signal - here we coulc specify further logic if we ever want to use co-op only.
      AtCo[i] = 0.36   if signals.get('signal') == 1 else (0.585  if signals.get('signal') == 0 else 0.055)
      AtPu[i] = 0.09   if signals.get('signal') == 1 else (0.47   if signals.get('signal') == 0 else 0.44)
      DiCo[i] = 0.045  if signals.get('signal') == 1 else (0.9275 if signals.get('signal') == 0 else 0.0275)
      DiPu[i] = 0.0225 if signals.get('signal') == 1 else (0.9225 if signals.get('signal') == 0 else 0.055)
      i+=1

   PE = numpy.prod(AtCo) + numpy.prod(AtPu) + numpy.prod(DiCo) + numpy.prod(DiPu)
   return [numpy.prod(AtCo)/PE, numpy.prod(AtPu)/PE, numpy.prod(DiCo)/PE, numpy.prod(DiPu)/PE]


#------ Spatial Functions ------#
def withinCOs(minCO, tCO, maxCO, inclusive=True): # Returns True if tCO lies within minCO and maxCO
    if inclusive:
        if tCO[0] < minCO[0]: return False
        if tCO[0] > maxCO[0]: return False
        if tCO[1] < minCO[1]: return False
        if tCO[1] > maxCO[1]: return False
    else:
        if tCO[0] <= minCO[0]: return False
        if tCO[0] >= maxCO[0]: return False
        if tCO[1] <= minCO[1]: return False
        if tCO[1] >= maxCO[1]: return False
    return True

def chunks(lst, n):
    """Yield successive n-sized chunks from lst."""
    for i in range(0, len(lst), n):
        yield lst[i:i + n]

def isLastMove(moveSet, currMove):
  if len(moveSet) == 0: return False
  if currMove == moveSet[-1]:
    return True
  else:
    return False
