import matplotlib.pyplot as p
import matplotlib,math,numpy,random

rel_error = 1e-15

# erf,erfc from http://www.digitalmars.com/archives/cplusplus/3634.html (bottom of page)

def erf(x):
    if abs(x) > 2.2:
        # use continued fraction for large arguments
        return 1.0 - erfc(x)
    sum = x
    term = x
    xsqr = x*x
    j = 1
    while True:
        term *= xsqr/j
        sum -= term/(2*j+1)
        j += 1
        term *= xsqr/j
        sum += term/(2*j+1)
        j += 1
        if sum == 0 or abs(term/sum) <= rel_error:
            break
    return (2.0/math.sqrt(math.pi)) * sum

def erfc(x):
    if abs(x) < 2.2:
        # use series when small arguments
        return 1.0 - erf(x)
    if x < 0:
        # continued fraction only valid for x > 0
        return 2.0 - erfc(-x)
    a = 1; b = x            # last two convergent numerators
    c = x; d = x*x + 0.5    # last two convergent denominators
    q1 = a/c; q2 = b/d      # last two convergents (a/c and b/d)
    n = 1.0
    while True:
        a,b = (b,a*n + b*x)
        c,d = (d,c*n + d*x)
        n += 0.5
        q1,q2 = (q2,b/d)
        if abs(q1-q2)/q2 <= rel_error:
            break
    return (1.0/math.sqrt(math.pi))*math.exp(-x*x)*q2

def unit_normal_cdf(x):
    if x < 0:
        return 1 - unit_normal_cdf(-x)
    else:
        return 0.5 + 0.5*erf(x/math.sqrt(2))

# convert digital sample sequence into voltages, upsample, add noise
# result is numpy array
def transmit(seq,vlow=0.0,vhigh=1.0,samples_per_bit=4,ntaps=0,bw=.08,nmag=0.0,nsigma=0.18):
    ndata = len(seq)*samples_per_bit
    samples = [vlow]*ntaps
    vlast = vlow
    for s in seq:
        v = vlow if s == 0 else vhigh
        dv = v - vlast
        for i in xrange(samples_per_bit):
            adjust = (dv*(i+1.1))/float(samples_per_bit) if i < samples_per_bit-1 else dv
            samples.append(adjust + vlast)
        vlast = v
    voltages = numpy.fromiter(samples,dtype=numpy.float)

    if ntaps > 0:
        taps = compute_taps(ntaps,bw)
        start = 3*ntaps/2
        filtered = numpy.convolve(voltages,taps)[start:start+ndata]
    else:
        filtered = voltages

    if nmag > 0:
        noise = nmag*numpy.random.normal(0,nsigma,ndata)
        """
        n3 = noise[3::4]
        f3 = filtered[3::4]
        print ndata,numpy.sum(f3 < 0.5),numpy.sum(f3 > 0.5),numpy.sum(numpy.abs(n3) > .5)
        err0 = numpy.sum(numpy.logical_and(f3 < 0.5,n3 > 0.5) * 1)
        err1 = numpy.sum(numpy.logical_and(f3 > 0.5,n3 < -0.5) * 1)
        print "errors for sample 3: ",err0+err1
        """
        return filtered+noise
    else:
        return filtered

# low-pass filter taps, cutoff is fraction of sample rate
def compute_taps(ntaps,cutoff,gain=1.0):
    order = float(ntaps - 1)
    # hamming window
    window = [0.53836 - 0.46164*numpy.cos((2*numpy.pi*i)/order)
              for i in xrange(ntaps)]

    fc = float(cutoff)
    wc = 2 * numpy.pi * fc
    middle = (ntaps - 1)/2
    taps = [0.0] * ntaps
    fmax = 0  # for low pass, gain @ DC = 1.0
    for i in xrange(ntaps):
        if i == middle:
            coeff = (wc/numpy.pi) * window[i]
            fmax += coeff
        else:
            n = i - middle
            coeff = (numpy.sin(n*wc)/(n*numpy.pi)) * window[i]
            fmax += coeff
        taps[i] = coeff
    gain = gain / fmax
    for i in xrange(ntaps):
        taps[i] *= gain
    return taps

def plot_eye_diagram(samples,samples_per_bit=4):
    p.figure()
    start = 0
    stop = len(samples) - samples_per_bit
    while start < stop:
        p.plot(samples[start:start+2*samples_per_bit+1])
        start += samples_per_bit
    p.title('Eye diagram')
    p.xlabel('Sample number')
    p.ylabel('volts')
    vmin = min(samples)
    vmax = max(samples)
    dv = vmax - vmin
    p.axis([0,2*samples_per_bit,vmin - 0.1*dv,vmax + 0.1*dv])

def plot_sample_histograms(samples,samples_per_bit=4,title='Sample distribution',nbins=100):
    minv = numpy.min(samples)
    maxv = numpy.max(samples)
    bins = numpy.reshape(samples,(-1,samples_per_bit))
    histograms = [numpy.histogram(bins[:,i],bins=nbins,range=(minv,maxv),new=True)[0]
                  for i in xrange(samples_per_bit)]
    nsamples = float(len(samples))/samples_per_bit
    maxh = max([max(histograms[i]) for i in xrange(samples_per_bit)])/nsamples
    for i in xrange(samples_per_bit):
        histograms[i] = histograms[i]/nsamples

    y = numpy.arange(minv,maxv,(maxv-minv)/float(nbins))
    p.figure()
    p.subplots_adjust(hspace=0.6)
    for i in xrange(samples_per_bit):
        p.subplot(1,samples_per_bit,i+1)
        p.hlines(y,[0],histograms[i],lw=2)
        p.title(str(i))
        p.axis([0,maxh,-.5,1.5])

random.seed(2536038)
message = [random.randint(0,1) for i in xrange(100000)]
samples_per_bit = 4
channel_data = transmit(message,samples_per_bit=samples_per_bit)
noisy_channel_data = transmit(message,samples_per_bit=samples_per_bit,nmag=1.0)

def verify_task1(sample_stats):
    stats = sample_stats(channel_data)
    expect = ((0.225,0.363,0.151),(0.025,0.263,0.126),(0.275,0.388,0.163),(0.500,0.500,0.250))
    error = False
    for i in xrange(samples_per_bit):
        smin,savg,savgsq = stats[i]
        emin,eavg,eavgsq = expect[i]
        if abs(smin - emin) > .001:
            print 'For sample time',i,'expected a min_dist of',emin,'but got',smin
            error = True
        if abs(savg - eavg) > .001:
            print 'For sample time',i,'expected an avg_dist of',eavg,'but got',savg
            error = True
        if abs(savgsq - eavgsq) > .001:
            print 'For sample time',i,'expected a avg_squared_dist of',eavgsq,'but got',savgsq
            error = True
    return None if error else 1

def verify_task4(bit_error_rate):
    message = [random.randint(0,1) for i in xrange(30)]
    xmessage = message[:]
    indices = (0,2,3,7,11,23,24,25,29)
    for index in indices:
        xmessage[index] = 1-message[index]
    ber = bit_error_rate(numpy.array(message),numpy.array(xmessage))
    eber = float(len(indices))/len(message)
    if ber != eber:
        print 'Called bit_error_rate, expected',eber,'got',ber
        print 'First arg: ',message
        print 'Second arg:',xmessage
        return None
    return 2

##################################################
##
## Code to submit task to server.  Do not change.
## Task-specific code is in verify(), defined above.
##
##################################################

import Tkinter
class Dialog(Tkinter.Toplevel):
    def __init__(self, parent, title = None):
        Tkinter.Toplevel.__init__(self, parent)
        self.transient(parent)
        if title: self.title(title)
        self.parent = parent

        body = Tkinter.Frame(self)
        self.initial_focus = self.body(body)
        body.pack(padx=5, pady=5)

        self.buttonbox()
        self.grab_set()

        if not self.initial_focus:
            self.initial_focus = self

        self.protocol("WM_DELETE_WINDOW", self.cancel)
        self.geometry("+%d+%d" % (parent.winfo_rootx()+50,parent.winfo_rooty()+50))
        
        self.initial_focus.focus_set()
        self.wait_window(self)

    def body(self, master):
        return None

    # add standard button box
    def buttonbox(self):
        box = Tkinter.Frame(self)
        w = Tkinter.Button(box, text="Ok", width=10, command=self.ok, default=Tkinter.ACTIVE)
        w.pack(side=Tkinter.LEFT, padx=5, pady=5)
        box.pack()
        
    # standard button semantics
    def ok(self, event=None):
        if not self.validate():
            self.initial_focus.focus_set() # put focus back
            return
        self.withdraw()
        self.update_idletasks()
        self.apply()
        self.cancel()
        
    def cancel(self, event=None):
        # put focus back to the parent window
        self.parent.focus_set()
        self.destroy()
        
    # command hooks
    def validate(self):
        return 1 # override

    def apply(self):
        pass   # override

# ask user for Athena username and MIT ID
class SubmitDialog(Dialog):
    def __init__(self,parent,error=None,title = None):
        self.error = error
        self.athena_name = None
        self.mit_id = None
        Dialog.__init__(self,parent,title=title)

    def body(self, master):
        row = 0
        if self.error:
            l = Tkinter.Label(master,text=self.error,
                              anchor=Tkinter.W,justify=Tkinter.LEFT,fg="red")
            l.grid(row=row,sticky=Tkinter.W,columnspan=2)
            row += 1
        Tkinter.Label(master, text="Athena username:").grid(row=row,sticky=Tkinter.E)
        self.e1 = Tkinter.Entry(master)
        self.e1.grid(row=row, column=1)

        row += 1
        Tkinter.Label(master, text="MIT ID:").grid(row=row,sticky=Tkinter.E)
        self.e2 = Tkinter.Entry(master)
        self.e2.grid(row=row, column=1)

        return self.e1 # initial focus

    # add standard button box
    def buttonbox(self):
        box = Tkinter.Frame(self)
        w = Tkinter.Button(box, text="Submit", width=10, command=self.ok,
                           default=Tkinter.ACTIVE)
        w.pack(side=Tkinter.LEFT, padx=5, pady=5)
        w = Tkinter.Button(box, text="Cancel", width=10, command=self.cancel)
        w.pack(side=Tkinter.LEFT, padx=5, pady=5)
        box.pack()
        
    def apply(self):
        self.athena_name = self.e1.get()
        self.mit_id = self.e2.get()

# Let user know what server said
class MessageDialog(Dialog):
    def __init__(self, parent,message = '',title = None):
        self.message = message
        Dialog.__init__(self,parent,title=title)

    def body(self, master):
        l = Tkinter.Label(master, text=self.message,anchor=Tkinter.W,justify=Tkinter.LEFT)
        l.grid(row=0)

# return contents of file as a string
def file_contents(fname):
    # use universal mode to ensure cross-platform consistency in hash
    f = open(fname,'U')
    result = f.read()
    f.close()
    return result

import hashlib
def digest(s):
    m = hashlib.md5()
    m.update(s)
    return m.hexdigest()

# if verify(f) indicates points have been earned, submit results
# to server if requested to do so
import inspect,os,urllib,urllib2
def checkoff(f,task='???',submit=True):
    if task == 'L3_1':
        points = verify_task1(f)
    elif task == 'L3_4':
        points = verify_task4(f)
    elif task == 'L3_6':
        points = 'tba'
    else:
        raise ValueError,"task must be one of L3_1, L3_4, or L3_6"

    if submit and points:
        root = Tkinter.Tk(); #root.withdraw()
        error = None
        while submit:
            sd = SubmitDialog(root,error=error,title="Submit Task %s?"%task)
            if sd.athena_name:
                if isinstance(f,str): fname = os.path.abspath(f)
                else: fname = os.path.abspath(inspect.getsourcefile(f))
                post = {
                    'user': sd.athena_name,
                    'id': sd.mit_id,
                    'task': task,
                    'digest': digest(file_contents(os.path.abspath(inspect.getsourcefile(checkoff)))),
                    'points': points,
                    'filename': fname,
                    'file': file_contents(fname)
                    }
                try:
                    response = urllib2.urlopen('http://scripts.mit.edu/~6.02/currentsemester/submit_task.cgi',
                                               urllib.urlencode(post)).read()
                except Exception,e:
                    response = 'Error\n'+str(e)
                if response.startswith('Error\n'):
                    error = response[6:]
                else:
                    MessageDialog(root,message=response,title='Submission response')
                    break
            else: break

        root.destroy()
