module ExitSwap

using MAT, JLD
using Grid
using NLopt

const nDim = 3

const vErrTol = 1e-5
const qErrTol = 1e-5
const maxIter = 1000
const pricesUpdateWeight = 0.10

type Equilibrium
	V::Array{Float64, nDim}
	Vd::Array{Float64, nDim}
	Vp::Array{Float64, nDim}
	Vaut::Array{Float64, 1}

	d::Array{Int, nDim}
	bPrS::Array{Float64, nDim}
	bPrL::Array{Float64, nDim}
	gammaS::Array{Float64, nDim}
	gammaL::Array{Float64, nDim}

	qS::Array{Float64, nDim}
	qL::Array{Float64, nDim}
	xiS::Array{Float64, nDim}
	xiL::Array{Float64, nDim}
	qRfS::Float64
	qRfL::Float64
end

type Calibration
	r::Float64
	eta::Float64
	lambda0::Float64
	lambda1::Float64
	beta::Float64
	deltaS::Float64
	deltaL::Float64
	alpha::Float64
	muS::Float64
end

type Grids
	bSgridSize::Int64
	bLgridSize::Int64
	bSgrid::FloatRange{Float64}
	bLgrid::FloatRange{Float64}
	yGridSize::Int64
	y::FloatRange{Float64}
	yPi::Array{Float64, 2}
end

type ModelStatistics
end

u(c) = -1 ./ c

h(yy, p::Calibration) = exp(yy) - max(0, p.lambda0 .* exp(yy) + p.lambda1 .* exp(yy).^2)

function findEquilibrium(g::Grids, p::Calibration)
	qRfS = 1 / (p.r - p.deltaS)
	qRfL = 1 / (p.r - p.deltaL)
	println("qRfS = $qRfS, qRfL = $qRfL")

	Vaut = (eye(g.yGridSize) - p.beta .* g.yPi) \ u(h(g.y, p))

	V0 = zeros(Float64, g.yGridSize, g.bSgridSize, g.bLgridSize)
	V1 = zeros(V0)
	Vd0 = zeros(V0)
	Vd1 = zeros(V0)
	Vp = zeros(V0)
	d = zeros(Int, g.yGridSize, g.bSgridSize, g.bLgridSize)
	bPrS = zeros(V0)
	bPrL = zeros(V0)
	gammaS = zeros(V0)
	gammaL = zeros(V0)
	qS = ones(V0) .* qRfS
	qL = ones(V0) .* qRfL
	xiS = ones(V0) .* qRfS ./ (1 + p.r)
	xiL = ones(V0) .* qRfL ./ (1 + p.r)

	for yIx = 1:g.yGridSize
		Vd0[yIx, :, :] = Vaut[yIx]
		V0[yIx, :, :] = Vaut[yIx]
	end

	println(nprocs(), " workers.")
	println("Start iterations...")
	vErr = 1.0
	qErr = 1.0
	iter = 0
	while iter < maxIter && (vErr > vErrTol || qErr > qErrTol)
		tic()
		iter += 1
		print("Iteration $iter: ")

		V0 = copy(V1)
		Vd0 = copy(Vd1)
	
		ranges = (g.y, g.bSgrid, g.bLgrid)
		V0f = CoordInterpGrid( ranges, V0, BCnearest, InterpLinear)
		qSf = CoordInterpGrid( ranges, qS, BCnearest, InterpLinear)
		qLf = CoordInterpGrid( ranges, qL, BCnearest, InterpLinear)
		# xiSf = CoordInterpGrid( ranges, xiS, BCnearest, InterpLinear)
		# xiLf = CoordInterpGrid( ranges, xiL, BCnearest, InterpLinear)
		# gammaSf = CoordInterpGrid( ranges, gammaS, BCnearest, InterpLinear)
		# gammaLf = CoordInterpGrid( ranges, gammaL, BCnearest, InterpLinear)

		function workAt(ix::Array{Int64, 1})
			yIx = ix[1]
			bSix = ix[2]
			bLix = ix[3]

			dPol = 0
			bSpol = 0.0
			bLpol = 0.0

			VpVal = 0.0
			Vval = 0.0

			VdVal = u(h(g.y[yIx], p)) + p.beta * sum( [ g.yPi[yIx,yPrIx] * ( 
				(1-p.eta) * Vd0[yPrIx, bSix, bLix] + p.eta * V0f[g.y[yPrIx], gammaS[yPrIx, bSix, bLix], gammaL[yPrIx, bSix, bLix]] ) for yPrIx = 1:g.yGridSize ])

			possibleOpt = Opt(:LN_SBPLX, 2)
			# maxeval!(possibleOpt, 5000)
			# maxtime!(possibleOpt, 60)
			lower_bounds!(possibleOpt, [g.bSgrid[1], g.bLgrid[1]])
			upper_bounds!(possibleOpt, [g.bSgrid[end], g.bLgrid[end]])

			revenuePossib = [ qS[yIx, bSPrIx, bLPrIx] * (g.bSgrid[bSPrIx] - (1 + p.deltaS) * g.bSgrid[bSix]) + qL[yIx, bSPrIx, bLPrIx] * (g.bLgrid[bLPrIx] - 
				(1 + p.deltaL) * g.bLgrid[bLix]) for bSPrIx = 1:g.bSgridSize, bLPrIx = 1:g.bLgridSize ]
			revenuePossibf = CoordInterpGrid( (g.bSgrid, g.bLgrid), revenuePossib, BCnearest, InterpLinear ) 

			max_objective!(possibleOpt, (x::Vector, g::Vector) -> revenuePossibf[x[1], x[2]])
			xtol_abs!(possibleOpt, 1e-6)
			(maxRev, bsAtMax, ret) = optimize!(possibleOpt, [g.bSgrid[bSix], g.bLgrid[bLix]])
			if ret != :XTOL_REACHED && ret != :ROUNDOFF_LIMITED
				println("maxRev ret ", ret, " at ", ix)
			end

			if (ret != :XTOL_REACHED && ret != :ROUNDOFF_LIMITED) || maxRev + exp(g.y[yIx]) - g.bSgrid[bSix] - g.bLgrid[bLix] <= 0
				# positive consumption is not possible
				VpVal = NaN
				bSpol = 0.0
				bLpol = 0.0
			else
				opt = Opt(:LN_SBPLX, 2)
				# maxeval!(opt, 50000)
				# maxtime!(opt, 60)
				lower_bounds!(opt, [g.bSgrid[1], g.bLgrid[1]])
				upper_bounds!(opt, [g.bSgrid[end], g.bLgrid[end]])

				expContVal = [ (g.yPi[yIx, :] * V0[:, bSPrIx, bLPrIx])[1] for bSPrIx = 1:g.bSgridSize, bLPrIx = 1:g.bLgridSize ]
				expContValf =  CoordInterpGrid( (g.bSgrid, g.bLgrid), expContVal, BCnearest, InterpLinear)

				function myObj(bbS::Float64, bbL::Float64)
					ctemp = exp(g.y[yIx]) - g.bSgrid[bSix] - g.bLgrid[bLix] + revenuePossibf[bbS, bbL]
					if ctemp <= 0
						println("Negative c = ", ctemp, " at ", bSix, " ", g.bSgrid[bSix], ", ",  bLix, " ", 
							g.bLgrid[bLix], " with possib. revenue ", revenuePossibf[bbS, bbL])
					end
					return ctemp <= 0 ? NaN : u(ctemp) + p.beta * expContValf[bbS, bbL]
				end

				#= myConst(bbS::Float64, bbL::Float64) = -(exp(g.y[yIx]) - g.bSgrid[bSix] - g.bLgrid[bLix] + revenuePossibf[bbS, bbL])
				inequality_constraint!(opt, (x::Vector, grad::Vector) -> myConst(x[1], x[2])) =#
								
				xtol_abs!(opt, 1e-6)
				max_objective!(opt, (x::Vector, grad::Vector) -> myObj(x[1], x[2]) )
				(maxVal, maxB, ret) = optimize!(opt, bsAtMax)
				if ret != :XTOL_REACHED && ret != :ROUNDOFF_LIMITED
					println("primes ret ", ret, " at ", ix)
				end
				VpVal = maxVal
				bSpol = maxB[1]
				bLpol = maxB[2]
			end

			if abs(g.bSgrid[bSix]) < 1e-10 && abs(g.bLgrid[bLix]) < 1e-10
				# no division by 0 in xi denominator
				xiSval = 0.0
				xiLval = 0.0
			else
				lendersSurplus = sum( [ g.yPi[yIx,yPrIx] * ( qSf[g.y[yPrIx], gammaS[yPrIx, bSix, bLix], gammaL[yPrIx, bSix, bLix]] * gammaS[yPrIx, bSix, bLix] + 
					qLf[g.y[yPrIx], gammaS[yPrIx, bSix, bLix], gammaL[yPrIx, bSix, bLix]] * gammaL[yPrIx, bSix, bLix] ) for yPrIx = 1:g.yGridSize ] )
					
				xiSval = ( p.eta * ( g.yPi[yIx,:] * xiS[:, bSix, bLix] )[1] + (1 - p.eta) * p.muS / (p.muS * g.bSgrid[bSix] + g.bLgrid[bLix]) * lendersSurplus ) / (1 + p.r)
				
				xiLval = ( p.eta * ( g.yPi[yIx,:] * xiL[:, bSix, bLix] )[1] + (1 - p.eta) / (p.muS * g.bSgrid[bSix] + g.bLgrid[bLix]) * lendersSurplus ) / (1 + p.r)
			end

			qSval = sum( [ g.yPi[yIx,yPrIx] * ( d[yPrIx, bSix, bLix] * xiS[yPrIx, bSix, bLix] + (1-d[yPrIx, bSix, bLix]) * (
				1 + (1 + p.deltaS) * qSf[g.y[yPrIx], bPrS[yPrIx, bSix, bLix], bPrL[yPrIx, bSix, bLix]] )) for yPrIx = 1:g.yGridSize ]) / (1 + p.r)

			qLval = sum( [ g.yPi[yIx,yPrIx] * ( d[yPrIx, bSix, bLix] * xiL[yPrIx, bSix, bLix] + (1-d[yPrIx, bSix, bLix]) * (
				1 + (1 + p.deltaL) * qLf[g.y[yPrIx], bPrS[yPrIx, bSix, bLix], bPrL[yPrIx, bSix, bLix]] )) for yPrIx = 1:g.yGridSize ]) / (1 + p.r)

			return VdVal, VpVal, bSpol, bLpol, xiSval, xiLval, qSval, qLval
		end

		indices = [ [yIx, bSix, bLix] for yIx = 1:g.yGridSize, bSix = 1:g.bSgridSize, bLix = 1:g.bLgridSize ]
		tmp = pmap(workAt, indices; err_retry = false)
		print("Primes map ", nprocs(), " ")
		tmp = reshape(tmp, size(V1))
		xiS1 = zeros(xiS)
		xiL1 = zeros(xiL)
		qS1 = zeros(qS)
		qL1 = zeros(qL)
		for yIx = 1:g.yGridSize, bSix = 1:g.bSgridSize, bLix = 1:g.bLgridSize
			Vd1[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][1]
			Vp[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][2]
			bPrS[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][3]
			bPrL[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][4]
			xiS1[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][5]	
			xiL1[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][6]
			qS1[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][7]	
			qL1[yIx, bSix, bLix] = tmp[yIx, bSix, bLix][8]

			d[yIx, bSix, bLix] = int( isnan(Vp[yIx, bSix, bLix]) || Vp[yIx, bSix, bLix] < Vd1[yIx, bSix, bLix] )
			V1[yIx, bSix, bLix] = d[yIx, bSix, bLix] == 1 ? Vd1[yIx, bSix, bLix] : Vp[yIx, bSix, bLix]
		end

		function gammaNash(yIx::Int)
			nashOpt = Opt(:LN_SBPLX, 2)
			# maxeval!(nashOpt, 5000)
			lower_bounds!(nashOpt, [g.bSgrid[1], g.bLgrid[1]])
			upper_bounds!(nashOpt, [g.bSgrid[end], g.bLgrid[end]])

			lenderValue = [ qS[yIx, bSPrIx, bLPrIx] * g.bSgrid[bSPrIx] + qL[yIx, bSPrIx, bLPrIx] * g.bLgrid[bLPrIx] for bSPrIx = 1:g.bSgridSize, bLPrIx = 1:g.bLgridSize ]
			lenderValuef = CoordInterpGrid( (g.bSgrid, g.bLgrid), lenderValue, BCnearest, InterpLinear )

			maxNash(bbS, bbL) = ( V0f[g.y[yIx], bbS, bbL] - Vaut[yIx] )^p.alpha * ( lenderValuef[bbS, bbL]  )^(1-p.alpha)
			max_objective!(nashOpt, (x::Vector, g::Vector) -> maxNash(x[1], x[2]))
			# inequality_constraint!(nashOpt, (x::Vector, grad::Vector) -> -( V0f[g.y[yIx], x[1], x[2]] - Vaut[yIx] ) , 1e-8)
			# inequality_constraint!(nashOpt, (x::Vector, grad::Vector) -> -( lenderValuef[x[1], x[2]] ), 1e-8)
			xtol_abs!(nashOpt, 1e-6)
			(maxN, bNash, ret) = optimize!(nashOpt, [1e-2, 1e-2])
			if ret != :XTOL_REACHED && ret != :ROUNDOFF_LIMITED
				println("Nash ret: ", ret)
			end
			return bNash
		end

		tmp2 = pmap(gammaNash, [ 1:g.yGridSize ]; err_retry = false)
		print("Nash map. ")
		for yIx = 1:g.yGridSize
			gammaS[yIx, :, :] = tmp2[yIx][1]
			gammaL[yIx, :, :] = tmp2[yIx][2]
		end

		vErr = (mean(abs(V1[:] - V0[:]).^2)).^(0.5)
		print("vErr = $vErr: ")
	
		qErr = (0.5 * mean(abs(qS1[:] - qS[:]).^2) + 0.5 * mean(abs(qL1[:] - qL[:]).^2)).^(0.5)
		print("qErr = $qErr: ")

		meanD = mean(d[:])
		print("mean d = $meanD: ")
		
		xiS = (1 - pricesUpdateWeight) * xiS + pricesUpdateWeight * xiS1
		xiL = (1 - pricesUpdateWeight) * xiL + pricesUpdateWeight * xiL1
		qS = qS1 # (1 - pricesUpdateWeight) * qS + pricesUpdateWeight * qS1
		qL = qL1 # (1 - pricesUpdateWeight) * qL + pricesUpdateWeight * qL1
		toc()
	end

	Equilibrium(V1, Vd1, Vp, Vaut, d, bPrS, bPrL, gammaS, gammaL, qS, qL, xiS, xiL, qRfS, qRfL)
end

function simulateEquilibrium(eq::Equilibrium, g::Grids, param::Calibration)
	ModelStatistics()
end

function portfolioDelta(bS::Float64, bL::Float64, deltaS::Float64, deltaL::Float64, t::Int)
	tt = [ 0:(t-1) ]
	lstream = log( bS .* (1 + deltaS).^tt + bL .* (1 + deltaL).^tt)
	X = [ones(t) tt]
	coef = (X' * X) \ X' * lstream
	return coef[2]
end

function compute()
	file = MAT.matopen("./data/endowment.mat")
	yData = MAT.read(file, "y")
	yGridSize = size(yData, 1)
	println("yGridSize = ", yGridSize)
	y = linrange(yData[1], yData[end], yGridSize)
	yPi = MAT.read(file, "yPi")
	MAT.close(file)

	bSmin = 0.0
	bSmax = 0.10
	bLmin = 0.0
	bLmax = 0.05
	bSgridSize = 30
	bLgridSize = 30
	println("bS grid size ", bSgridSize, " in [", bSmin, ", ", bSmax, "]")
	println("bL grid size ", bLgridSize, " in [", bLmin, ", ", bLmax, "]")
	bSgrid = linrange(bSmin, bSmax, bSgridSize)
	bLgrid = linrange(bLmin, bLmax, bLgridSize)
	g = Grids(bSgridSize, bLgridSize, bSgrid, bLgrid, yGridSize, y, yPi)

	r = 0.01
	eta = 0.0385
	lambda0 = -0.15
	lambda1 = 0.155
	beta = 0.94
	deltaS = -0.10
	deltaL = -0.03
	alpha = 0.25
	muS = 0.33
	param = Calibration(r, eta, lambda0, lambda1, beta, deltaS, deltaL, alpha, muS)	

	eq = findEquilibrium(g, param)
	@save "equilibrium.jld" eq g param
	stat = simulateEquilibrium(eq, g, param)
	println("Computation done.")
end

end
