import numpy as np
import sys
import os
import random
from random import seed
from random import random
from random import gauss
import datetime
import matplotlib.pyplot as plt
from matplotlib.backends.backend_pdf import PdfPages
from matplotlib.ticker import MultipleLocator
from matplotlib.ticker import ScalarFormatter
from matplotlib.ticker import NullFormatter
import matplotlib.image as image
import matplotlib.patches as patches
import matplotlib.lines as lines
import matplotlib.ticker as ticker
from scipy.optimize import curve_fit

from matplotlib.patches import Circle, Ellipse, Wedge, Polygon
from matplotlib.collections import PatchCollection
import matplotlib.path as mpath
import matplotlib.patches as mpatches
from matplotlib.patches import FancyArrowPatch

Path = mpath.Path



def global_variables ():
	# variabili comuni utili per le varie figure
	global axisticslabelfontsize
	global axisticslabelfontsizeinset
	global axislabelfontsize 
	global axislabelfontsizeinset
	global legendfontsize
	global linewidth
	global linewidth2
	global linewidth3
	global linewidth4
	global markersize
	# colors (colorblind safe)
	global color1
	global color2
	global color3
	
	axisticslabelfontsize=9
	axisticslabelfontsizeinset=7
	axislabelfontsize=11 
	axislabelfontsizeinset=8
	legendfontsize=10
	linewidth = 2.0
	linewidth2 = 3.0
	linewidth3 = 1.5
	linewidth4 = 6.0
	markersize = 1.0
	color1 = "#d95f02"
	color2 = "#1b9e77"
	color3 = "#7570b3"
	
	return

# create a mock alternating ABP/BP trajectory
def generate_mock_trajectory(sigma,v,D,Drot,Dt):
	switch_prob = [0.01,0.001]
	x, y =  3,-4 # start outside target
	factor = np.sqrt(2*D*Dt)
	factor2 = np.sqrt(2*Drot*Dt)
	
	phases = []
	pos = []
	l = 1
	nt = 0
	phase = 1
	while (nt<10000):
		theta = np.random.uniform(0, 2*np.pi)
		phase +=1
		if (phase==2): phase = 0
		
		n_steps  = 1
		while(np.random.uniform(0, 1) > switch_prob[phase]):
			n_steps+=1
		
		xs=[]
		ys=[]
		for t in range(n_steps):
			nt += 1
			if (nt>10000): break
			x += phase*v*Dt*np.cos(theta) + factor*np.random.normal(0, 1.0)
			y += phase*v*Dt*np.sin(theta) + factor*np.random.normal(0, 1.0)
			theta += factor2 *np.random.normal(0, 1.0)
			xs.append(x)
			ys.append(y)
			
			if (phase==0):
				r = np.sqrt(x*x+y*y)
				if (r<=sigma/2):
					l=0
					break;
						
		phases.append(phase)
		pos.append([xs,ys])
		
		if (l==0): break
		
	return phases,pos,l

def generate_mock_trajectory_2(sigma,v,D,Drot,Dt):
	switch_prob = [0.7,0.1]
	x, y =  5,0.1 # start outside target
	factor = np.sqrt(2*D*Dt)
	factor2 = np.sqrt(2*Drot*Dt)
	
	phases = []
	pos = []
	phase = 1
	thetas = [np.pi,0.75*np.pi,np.pi,0.35*np.pi,np.pi]
	for nt in range(5):
		theta = thetas[nt]
		phase +=1
		if (phase==2): phase = 0
		
		n_steps=1
		while(np.random.uniform(0, 1) > switch_prob[phase]):
			n_steps+=1
		
		xs=[x]
		ys=[y]
		for t in range(n_steps):
			x += phase*v*Dt*np.cos(theta) + factor*np.random.normal(0, 1.0)
			y += phase*v*Dt*np.sin(theta) + factor*np.random.normal(0, 1.0)
			theta += factor2 *np.random.normal(0, 1.0)
			xs.append(x)
			ys.append(y)
						
		phases.append(phase)
		pos.append([xs,ys])
		
	return phases,pos


def make_sketch (sigma,Rtilde,phases,pos,phases2,pos2):
	xfig,yfig=7.0,2.0
	factor = xfig/yfig
	
	x1,y1 = 0,0
	dx1,dy1=1.0,1.0
	
	x2,y2 = 0.05,0.05
	dy2 = 0.9
	dx2 = dy2 *  yfig/xfig
	
	x3,y3 = 0.38,0.05
	dy3 = 0.9
	dx3 = dy3 *  yfig/xfig
	
	x4,y4 = 0.72,0.05
	dy4 = 0.9
	dx4 = dy4 *  yfig/xfig
	

	with PdfPages('fig1.pdf') as pdf:
		fig = plt.figure(figsize=(xfig,yfig))
		plt.rc('text', usetex=True)
		plt.rc('text.latex', preamble = ','.join('''\usepackage{txfonts} \usepackage{lmodern}'''.split()))
		
		######################## panel main ########################################
		panel = fig.add_axes([x1, y1, dx1, dy1])
		panel.set_xticks([])
		panel.set_yticks([])
		panel.set_xticklabels([])
		panel.set_yticklabels([])
		
		panel.text(0.01,0.9,r'{\bf (a)}',fontsize=legendfontsize,transform=panel.transAxes)
		panel.text(0.34,0.9,r'{\bf (b)}',fontsize=legendfontsize,transform=panel.transAxes)
		panel.text(0.68,0.9,r'{\bf (c)}',fontsize=legendfontsize,transform=panel.transAxes)
		
		
		
		######################## panel trajectory ########################################
		panel = fig.add_axes([x4, y4, dx4, dy4])
		panel.spines['right'].set_visible(False)
		panel.spines['top'].set_visible(False)
		panel.spines['bottom'].set_visible(False)
		panel.spines['left'].set_visible(False)
		panel.axes.get_xaxis().set_ticks([])
		panel.axes.get_yaxis().set_ticks([])
		
		DeltaX=14
		xlim_low=-3
		ylim_low=-10.1
		panel.set_xlim(xlim_low,xlim_low+DeltaX)
		panel.set_ylim(ylim_low,ylim_low+DeltaX)
		
		
		c_target = Circle((0.0,0.0),sigma/2,facecolor='lightgreen',edgecolor='k',lw=0.8)
		panel.add_patch(c_target)
		
		c_ext = Circle((0.0,0.0),Rtilde,facecolor='none',edgecolor='k',linestyle="--",lw=0.8)
		panel.add_patch(c_ext)
		
		
		x1,y1 = [pos[-1][0][-1],0.9*pos[-1][0][-1]],[pos[-1][1][-1],0.9*pos[-1][1][-1]]
		
		panel.plot(pos[-6][0],pos[-6][1],'r-',linewidth=0.8,label=r'ABP ($\phi=1$)')
		panel.plot(pos[-4][0],pos[-4][1],'r-',linewidth=0.8)
		panel.plot(pos[-2][0],pos[-2][1],'r-',linewidth=0.8)
		panel.plot(pos[-7][0],pos[-7][1],'b-',linewidth=1.2,label=r'BP ($\phi=0$)')
		panel.plot(pos[-5][0],pos[-5][1],'b-',linewidth=1.2)
		panel.plot(pos[-3][0],pos[-3][1],'b-',linewidth=1.2)
		panel.plot(pos[-1][0],pos[-1][1],'b-',linewidth=1.2)
		panel.plot(x1,y1,'b-',linewidth=1.2)
		
		psi = np.pi/3
		panel.arrow(sigma*np.cos(-(np.pi/2-psi)) -sigma*np.cos(psi)/2, sigma*np.sin(-(np.pi/2-psi))-sigma*np.sin(psi)/2,  sigma*np.cos(psi), sigma*np.sin(psi),  lw=0.5, head_width=0.1, color = 'gray', length_includes_head=True)
		panel.arrow(sigma*np.cos(-(np.pi/2-psi)) +sigma*np.cos(psi)/2, sigma*np.sin(-(np.pi/2-psi))+sigma*np.sin(psi)/2,  -sigma*np.cos(psi), -sigma*np.sin(psi),  lw=0.5, head_width=0.1, color = 'gray', length_includes_head=True)
		panel.text( 1.2*sigma*np.cos(-(np.pi/2-psi)), 1.2*sigma*np.sin(-(np.pi/2-psi))- 0.2,r'$\sigma$',fontsize=axisticslabelfontsize-1,color='gray')
		
		psi = np.pi/9
		panel.arrow(0, 0,  Rtilde*np.cos(psi), Rtilde*np.sin(psi),  lw=0.5, head_width=0.1, color = 'gray', length_includes_head=True)
		panel.arrow(Rtilde*np.cos(psi), Rtilde*np.sin(psi),  -Rtilde*np.cos(psi), -Rtilde*np.sin(psi),  lw=0.5, head_width=0.1, color = 'gray', length_includes_head=True)
		panel.text( 4, 2,r'$\tilde{R}$',fontsize=axisticslabelfontsize-1,color='gray')
		
		psi = np.pi/2.8
		panel.arrow(2.2*np.cos(psi), 2.2*np.sin(psi),  -1.5*np.cos(psi), -1.5*np.sin(psi),  lw=0.5, head_width=0.1, color = 'k', length_includes_head=True)
		panel.text( 0.4, 2.4,r'target',fontsize=axisticslabelfontsize-1,color='k')
		
		handles, labels = panel.get_legend_handles_labels()
		panel.legend(handles[::-1], labels[::-1],loc='upper right',bbox_to_anchor=(0.34, 0.25),ncol=1,fontsize=legendfontsize-3,handlelength=1.0,labelspacing=0.2,framealpha=1.0)
		
		
		######################## panel state ########################################
		panel = fig.add_axes([x2, y2, dx2, dy2])
		panel.spines['right'].set_visible(False)
		panel.spines['top'].set_visible(False)
		panel.spines['bottom'].set_visible(False)
		panel.spines['left'].set_visible(False)
		panel.axes.get_xaxis().set_ticks([])
		panel.axes.get_yaxis().set_ticks([])
		
		DeltaX=8
		xlim_low=-1
		ylim_low=-5
		panel.set_xlim(xlim_low,xlim_low+DeltaX)
		panel.set_ylim(ylim_low,ylim_low+DeltaX)
		
		
		panel.plot(pos2[-4][0],pos2[-4][1],'r-',linewidth=1,label=r'ABP ($\phi=1$)')
		panel.plot(pos2[-2][0],pos2[-2][1],'r-',linewidth=1)
		panel.plot(pos2[-5][0],pos2[-5][1],'b-',linewidth=1.4,label=r'BP ($\phi=0$)')
		panel.plot(pos2[-3][0],pos2[-3][1],'b-',linewidth=1.4)
		panel.plot(pos2[-1][0],pos2[-1][1],'b-',linewidth=1.4)
		
		handles, labels = panel.get_legend_handles_labels()
		panel.legend(handles[::-1], labels[::-1],loc='upper right',bbox_to_anchor=(1.02, 1.01),ncol=1,fontsize=legendfontsize-3,handlelength=1.0,labelspacing=0.2,framealpha=1.0)
		
		c_target = Circle((0.0,0.0),sigma/2,facecolor='lightgreen',edgecolor='k',lw=0.8)
		panel.add_patch(c_target)
		psi = 0.6*np.pi
		panel.arrow(1.3*np.cos(psi), 1.3*np.sin(psi),  -0.7*np.cos(psi), -0.7*np.sin(psi),  lw=0.5, head_width=0.1, color = 'k', length_includes_head=True)
		panel.text( -1.4, 1.4,r'target',fontsize=axisticslabelfontsize-1,color='k')
		

		panel.arrow(0, 0, pos2[-4][0][2] ,pos2[-4][1][2] ,  lw=0.5, head_width=0.1, color = 'gray', length_includes_head=True)
		panel.arrow(0, 0, pos2[-4][0][3] ,pos2[-4][1][3] ,  lw=0.5, head_width=0.1, color = 'gray', length_includes_head=True)
		panel.text( 2.3,0.05,r'$r_{t-\Delta t}$',fontsize=axisticslabelfontsize-1,color='gray')
		panel.text( 1.45,1.0,r'$r_t$',fontsize=axisticslabelfontsize-1,color='gray')
		
		
		text = r'$\omega_t = \! \left\lbrace \!\!\! \begin{array}{ll} 1  & \mbox{if } r_t \!<\! r_{t-\Delta t} \\ 0 & \mbox{otherwise} \end{array} \right. $'
		panel.text( 1.9, -1.4,text,fontsize=legendfontsize-3,color='k')
		
		
		panel.text( 0.4, -2.6,r'agent  $\quad$ state',fontsize=legendfontsize-3,color='k')
		l1 = lines.Line2D([0,6], [-2.8,-2.8], color='k', alpha=1.0, linewidth=0.4)
		panel.add_line(l1)
		l1 = lines.Line2D([1.9,1.9], [-2,-4.6], color='k', alpha=1.0, linewidth=0.4)
		panel.add_line(l1)
		panel.text( 0.4, -3.4,r'type A $\quad s_t = (\phi_t,r_t)$',fontsize=legendfontsize-3,color='k')
		panel.text( 0.4, -4,r'type B $\quad s_t = (\phi_t,\omega_t)$',fontsize=legendfontsize-3,color='k')
		panel.text( 0.4, -4.6,r'type C $\quad s_t = (\phi_t,r_t,\omega_t)$',fontsize=legendfontsize-3,color='k')
		
		
		
		######################## panel action ########################################
		panel = fig.add_axes([x3, y3, dx3, dy3])
		panel.spines['right'].set_visible(False)
		panel.spines['top'].set_visible(False)
		panel.spines['bottom'].set_visible(False)
		panel.spines['left'].set_visible(False)
		panel.axes.get_xaxis().set_ticks([])
		panel.axes.get_yaxis().set_ticks([])
		
		DeltaX=8
		xlim_low=-4
		ylim_low=-4
		panel.set_xlim(xlim_low,xlim_low+DeltaX)
		panel.set_ylim(ylim_low,ylim_low+DeltaX)
		
		
		text = r'$s_t \Rightarrow a_t  = \! \left\lbrace \!\!\!   \begin{array}{l} \mbox{keep phase } (\phi_{t+\Delta t} \!=\! \phi_t)   \\ \mbox{switch phase } (\phi_{t+\Delta t} \!=\! 1 \!-\! \phi_t)  \end{array}  \right. $'
		panel.text( -3.8, 3.0,text,fontsize=legendfontsize-3.5,color='k')
		
		
		
		xs=np.asarray([0,1.2,2.8])-3.2
		ys=np.asarray([0,0.4,1.5])-0.6
		x2s = [xs[-1],xs[-1]+0.6,xs[-1]+0.2,xs[-1]+0.5]
		y2s = [ys[-1],ys[-1]-0.5,ys[-1]-0.4,ys[-1]-0.2]
		x3s = [xs[-1],xs[-1]+1.2,xs[-1]+2.4,xs[-1]+3.2]
		y3s = [ys[-1],ys[-1]+0.5,ys[-1]+0.8,ys[-1]+0.6]
		
		
		panel.plot(xs[:2],ys[:2],'o-',color='gray',linewidth=1,markersize=1)
		panel.plot(xs[1:],ys[1:],'ro-',linewidth=1,markersize=1,label=r'ABP ($\phi=1$)')
		panel.plot(x2s,y2s,'bo--',linewidth=1,markersize=1,label=r'BP ($\phi=0$)')
		panel.plot(x3s,y3s,'ro--',linewidth=1,markersize=1,label=r'BP ($\phi=0$)')
		panel.plot(xs[-1],ys[-1],'ks',markersize=2)
		
		
		
		
		panel.text( -3.7, 0.3,r'$\phi_t \!=\! 1$',fontsize=legendfontsize-4,color='k',bbox=dict(facecolor='white',edgecolor='k', boxstyle='round,pad=0.3',lw=0.5,alpha=0.8))
		arrow = FancyArrowPatch((-3, 0.05), (xs[1]-0.02, ys[1]+0.02),
                        connectionstyle="arc3,rad=0.0",  # curvature
                        arrowstyle="-|>",                # simple arrow head
                        mutation_scale=3,               # size of arrow head
                        lw=0.5, color="k")
		panel.add_patch(arrow)
		
		panel.text( 1.22, -0.62,r'$a_t = \mathrm{switch}$' '\n' r'$\phi_{t+\Delta t} = 0$',fontsize=legendfontsize-4,color='b',ha='center',bbox=dict(facecolor='white',edgecolor='b', boxstyle='round,pad=0.3',lw=0.5,alpha=0.8))
		arrow = FancyArrowPatch((0, -0.3), (xs[2], ys[2]-0.03),
                        connectionstyle="arc3,rad=-0.7",  # curvature
                        arrowstyle="-|>",                # simple arrow head
                        mutation_scale=3,               # size of arrow head
                        lw=0.5, color="b")
		panel.add_patch(arrow)
		
		panel.text(-2.0, 1.8,r'$a_t = \mathrm{keep}$' '\n' r'$\phi_{t+\Delta t} = 1$',fontsize=legendfontsize-4,color='r',bbox=dict(facecolor='white',edgecolor='r', boxstyle='round,pad=0.3',lw=0.5,alpha=0.8))
		arrow = FancyArrowPatch((-0.3, 1.65), (xs[2], ys[2]+0.03),
                        connectionstyle="arc3,rad=0.2",  # curvature
                        arrowstyle="-|>",                # simple arrow head
                        mutation_scale=3,               # size of arrow head
                        lw=0.5, color="r")
		panel.add_patch(arrow)
		
		
		panel.text( 0, -1.6,'Stochastic policy:',fontsize=legendfontsize-3.5,color='k',ha="center")
		panel.text( 0, -2.1,'switch phase with probability $p_t$',fontsize=legendfontsize-3.5,color='k',ha="center")
		
		
		
		c_target = Circle((-0.5,-3.3),0.3,facecolor='lightgreen',edgecolor='k',lw=0.8)
		panel.add_patch(c_target)
		
		xs=np.asarray([-0.9,-0.6])
		ys=np.asarray([-3.2,-3.4])
		x2s=np.asarray([-3.2,-2.4,-1.8,xs[0]])
		y2s=np.asarray([-3.4,-3.3,-3,ys[0]])
		x3s=np.asarray([-3.4,-3.3,-3.1,x2s[0]])
		y3s=np.asarray([-3.35,-3.2,-3.1,y2s[0]])
		panel.plot(x3s,y3s,'bo-',linewidth=1,markersize=1)
		panel.plot(x2s,y2s,'ro-',linewidth=1,markersize=1)
		panel.plot(xs,ys,'bo-',linewidth=1,markersize=1)
		
		
		panel.text( 2, -2.9,'Positive reward',fontsize=legendfontsize-3.5,color='k',ha="center")
		panel.text( 2, -3.4,'if target is met',fontsize=legendfontsize-3.5,color='k',ha="center")
		panel.text( 2, -3.9,'when $\phi_t=0$',fontsize=legendfontsize-3.5,color='k',ha="center")
		
		
		pdf.savefig(fig)





def main():
	global_variables ()
	
	# setting units
	sigma = 1.0
	tau = 1.0
	D = sigma*sigma/ (4*tau)
	
	Pe = 100
	v = Pe * sigma/tau
	ell = 1
	Drot = v/ (ell*sigma)
	
	Rtilde = 10*sigma
	
	Dt = 0.0001
	
	
	n = 0
	l = 1
	np.random.seed(5)
	while(l==1):
		n += 1
		print 'generate trajectory',n
		phases,pos,l = generate_mock_trajectory(sigma,v,D,Drot,Dt)
	
	np.random.seed(5)
	phases2,pos2 = generate_mock_trajectory_2(sigma,v,D,Drot/100,0.008)
	
	make_sketch (sigma,Rtilde,phases,pos,phases2,pos2)





# Drot/2 rseed1



main()
