import numpy as np
import scipy as sp 
import matplotlib.pyplot as plt
import optimize
Nx = 9; Lx = 8.0
Ny = 5; Ly = 4.0

xn = np.linspace(0.0, Lx, Nx)
yn  = np.linspace(0.0, Ly, Ny)

X, Y = np.meshgrid(xn, yn)

nodes = np.vstack((np.reshape(X, (Nx*Ny,)), np.reshape(Y, (Nx*Ny,)))).T
NoofNodes = nodes.shape[0]

members = []
# add horizontal members
for j in range(0, Ny):
	for i in range(0, Nx-1):
		members.append([j*Nx + i, j*Nx + i +1])

for j in range(0, Ny-1):
	for i in range(1, Nx):
		members.append([j*Nx + i, (j+1)*Nx + i])

for j in range(0, Ny-1):
	for i in range(0, Nx-1):
		members.append([j*Nx + i, (j+1)*Nx + i+1])
		members.append([j*Nx + i+1, (j+1)*Nx + i])

members = np.array(members)
lengths = np.empty(members.shape[0])
xd = np.empty(members.shape[0])
xd[:] = 1.0e-4
Vmax = 0.0
# for m in range(0, members.shape[0]):
# 	line = nodes[members[m]]
# 	plt.plot(line[:,0], line[:,1], '-k', linewidth=2.0)
# plt.show()
for m in range(0, members.shape[0]):
	line = nodes[members[m]]
	lengths[m] = np.sqrt((line[1,0] - line[0,0])**2 + (line[1,1] - line[0,1])**2)
	Vmax += lengths[m]*xd[m]

print("initial/Max volume = ", Vmax)

E0 = 10000
F_G = np.zeros((2*NoofNodes, 1))
u_G = np.zeros((2*NoofNodes, 1))
K_G = np.zeros((2*NoofNodes, 2*NoofNodes))

nodes_mask = np.abs(nodes[:,0])>1.0e-5
nodes_mask = np.repeat(nodes_mask, 2)

iF = np.where((np.abs(nodes[:,0] - Lx)< 1.0e-8) & (np.abs(nodes[:,1])<1.0e-8))[0]
F_G[2*iF + 1] = -1.0

k0js = []
for m in range(0, members.shape[0]):
	line = nodes[members[m]]
	l = lengths[m]
	c = (line[1,0] - line[0,0])/l
	s = (line[1,1] - line[0,1])/l

	kj = np.array([
		[c*c, c*s, -c*c, -c*s],
		[c*s, s*s, -c*s, -s*s],
		[-c*c, -c*s, c*c, c*s],
		[-c*s, -s*s, c*s, s*s],
		])
	k0js.append(kj)

def solve_truss(x_in, nodes_mask):
	K_G[:] = 0.0
	for m in range(0, members.shape[0]):
		l = lengths[m]
		K_G[2*members[m][0]:2*members[m][0]+2, 2*members[m][0]:2*members[m][0]+2] += x_in[m]*E0*k0js[m][0:2, 0:2]/l

		K_G[2*members[m][0]:2*members[m][0]+2, 2*members[m][1]:2*members[m][1]+2] += x_in[m]*E0*k0js[m][0:2, 2:4]/l

		K_G[2*members[m][1]:2*members[m][1]+2, 2*members[m][0]:2*members[m][0]+2] += x_in[m]*E0*k0js[m][2:4, 0:2]/l

		K_G[2*members[m][1]:2*members[m][1]+2, 2*members[m][1]:2*members[m][1]+2] += x_in[m]*E0*k0js[m][2:4, 2:4]/l

	# plt.spy(K_G, precision=1.0e-6)
	# plt.show()
	KR = K_G[nodes_mask][:,nodes_mask]
	FR = F_G[nodes_mask]

	uR = sp.linalg.solve(KR, FR)
	u_G[nodes_mask] = uR
	C = np.sum(FR*uR) #FR.dot(uR)
	return u_G, C
def getxStar(lam, Lj, q0j, length, alpha):
	xt = Lj + np.sqrt(q0j/(lam*length))
	if(xt < alpha):
		return alpha
	elif(xt > xmax):
		return xmax
	else:
		return xt

def getDual(lam, r0, Vmax, q0s, Ljs, alphas, lengths, Noofels):
	phi = r0 - lam*Vmax
	for m in range(0, Noofels):
		xStar = getxStar(lam, Ljs[m], q0s[m], lengths[m], alphas[m])
		phi += q0s[m]/(xStar - Ljs[m]) + lam*lengths[m]*xStar
	return -phi
q0s = np.zeros(members.shape[0])
Ljs = np.zeros(members.shape[0])
Ljk_1 = np.zeros(members.shape[0])

alphas = np.zeros(members.shape[0])
xk = np.zeros(members.shape[0])
xk_1 = np.empty(members.shape[0])
xk_2 = np.empty(members.shape[0])
maxk = 200
k_opt = 0
xk[:] = xd[:]
xmax = 1.0e-2
xmin = 1.0e-6
s_init = 1.0e-2
s_faster = 1.15
s_slower = 0.75
mu = 0.15
lam_init = 1.0e-8
lam_star = lam_init
plotscale = 1000
while(k_opt < maxk):
	u_G, C = solve_truss(xk, nodes_mask)
	r0 = C
	if(k_opt < 2):
		Ljs[:] = xk[:] - s_init*(xmax - xmin)
	else:
		for m in range(0, members.shape[0]):
			if((xk[m] - xk_1[m])*(xk_1[m] - xk_2[m]) > 0.0):
				Ljs[m] = xk[m] - s_faster*(xk_1[m] - Ljk_1[m])
			else:
				Ljs[m] = xk[m] - s_slower*(xk_1[m] - Ljk_1[m])
	for m in range(0, members.shape[0]):
		alphas[m] = max(xmin, Ljs[m] + mu*(xk[m] - Ljs[m]))
		um = np.array([u_G[2*members[m][0]],u_G[2*members[m][0]+1], u_G[2*members[m][1]],u_G[2*members[m][1]+1]])
		# breakpoint()
		uTk0u = np.dot(um.T, np.dot(k0js[m], um))[0][0]
		q0s[m] = ((xk[m] - Ljs[m])**2)*uTk0u
		r0 += -(xk[m] - Ljs[m])*uTk0u
		lam_crit = q0s[m]/(lengths[m]*(xmax - Ljs[m])**2)
		# print(m, lam_crit)
	# goldenSection(f, a_in, b_in, tol, maxiter, *args)
	args_list = [ r0, Vmax, q0s, Ljs, alphas, lengths, members.shape[0]]
	lambda_opt = optimize.goldenSection(getDual, 0.0, 1.0, 1.0e-8, 1000, *args_list)
	lam_star = lambda_opt
	negphi_k = getDual(lam_star, *args_list)
	xk_2[:] = xk_1[:]
	xk_1[:] = xk[:]

	Ljk_1[:] = Ljs[:]

	for m in range(0, members.shape[0]):
		xk[m] = getxStar(lam_star, Ljs[m], q0s[m], lengths[m], alphas[m])
	# print(xk)
	print("saving iteration ", k_opt)
	fig = plt.figure(figsize=(9, 5))
	for m in range(0, members.shape[0]):
		line = nodes[members[m]]
		plt.plot(line[:,0], line[:,1], '-k', linewidth=plotscale*xk[m])

	plt.savefig("./trussIterationsClass/iter="+"%04d"%k_opt+".png")
	plt.close()
	k_opt = k_opt+1


	# negphis = []
	# lambdas = []
	# for j in range(0, 100):
	# 	lam_plot = lam_init + j*5.0e-3
	# 	lambdas.append(lam_plot )
	# 	negphis.append(getDual(lam_plot, r0, Vmax, q0s, Ljs, alphas, lengths, members.shape[0]))
	# plt.plot(lambdas, negphis)
	# plt.plot(lambda_opt, getDual(lambda_opt, *args_list), 'ko')
	# plt.show()
	# breakpoint()
# solve_truss(xd, nodes_mask)
breakpoint()