#!/usr/bin/env python
# coding: utf-8

# In[1]:


from sympy import symbols
import sympy.physics.mechanics as me


# In[2]:


n = 3


# In[3]:


# Generalized coordinates
alpha = me.dynamicsymbols('alpha:{}'.format(n))
beta = me.dynamicsymbols('beta:{}'.format(n))

# Generalized speeds
omega = me.dynamicsymbols('omega:{}'.format(n))
delta = me.dynamicsymbols('delta:{}'.format(n))


# In[4]:


m_bob = symbols('m:{}'.format(n))


# In[5]:


l = symbols('l:{}'.format(n))
m_link = symbols('M:{}'.format(n))
Ixx = symbols('Ixx:{}'.format(n))
Iyy = symbols('Iyy:{}'.format(n))
Izz = symbols('Izz:{}'.format(n))


# In[6]:


g = symbols('g')


# In[7]:


I = me.ReferenceFrame('I')


# In[8]:


A = me.ReferenceFrame('A')
A.orient(I, 'Space', [alpha[0], beta[0], 0], 'ZXY')

B = me.ReferenceFrame('B')
B.orient(A, 'Space', [alpha[1], beta[1], 0], 'ZXY')

C = me.ReferenceFrame('C')
C.orient(B, 'Space', [alpha[2], beta[2], 0], 'ZXY')


# In[9]:


kinematic_differentials = []
for i in range(n):
   kinematic_differentials.append(omega[i] - alpha[i].diff())
   kinematic_differentials.append(delta[i] - beta[i].diff())


# In[10]:


A.set_ang_vel(I, omega[0] * I.z + delta[0] * I.x)
B.set_ang_vel(I, omega[1] * I.z + delta[1] * I.x)
C.set_ang_vel(I, omega[2] * I.z + delta[2] * I.x)


# In[11]:


O = me.Point('O')
O.set_vel(I, 0)


# In[12]:


P1 = O.locatenew('P1', -l[0] * A.y)
P2 = P1.locatenew('P2', -l[1] * B.y)
P3 = P2.locatenew('P3', -l[2] * C.y)


# In[13]:


P1.v2pt_theory(O, I, A)
P2.v2pt_theory(P1, I, B)
P3.v2pt_theory(P2, I, C)
points = [P1, P2, P3]


# In[14]:


Pa1 = me.Particle('Pa1', points[0], m_bob[0])
Pa2 = me.Particle('Pa2', points[1], m_bob[1])
Pa3 = me.Particle('Pa3', points[2], m_bob[2])
particles = [Pa1, Pa2, Pa3]


# In[15]:


P_link1 = O.locatenew('P_link1', -l[0] / 2 * A.y)
P_link2 = P1.locatenew('P_link2', -l[1] / 2 * B.y)
P_link3 = P2.locatenew('P_link3', -l[2] / 2 * C.y)


# In[16]:


P_link1.v2pt_theory(O, I, A)
P_link2.v2pt_theory(P1, I, B)
P_link3.v2pt_theory(P2, I, C)

points_rigid_body = [P_link1, P_link2, P_link3]


# In[17]:


inertia_link1 = (me.inertia(A, Ixx[0], Iyy[0], Izz[0]), P_link1)
inertia_link2 = (me.inertia(B, Ixx[1], Iyy[1], Izz[1]), P_link2)
inertia_link3 = (me.inertia(C, Ixx[2], Iyy[2], Izz[2]), P_link3)


# In[18]:


link1 = me.RigidBody('link1', P_link1, A, m_link[0], inertia_link1)
link2 = me.RigidBody('link2', P_link2, B, m_link[1], inertia_link2)
link3 = me.RigidBody('link3', P_link3, C, m_link[2], inertia_link3)
links = [link1, link2, link3]


# In[19]:


forces = []

for particle in particles:
   mass = particle.mass
   point = particle.point
   forces.append((point, -mass * g * I.y))

for link in links:
   mass = link.mass
   point = link.masscenter
   forces.append((point, -mass * g * I.y))


# In[20]:


total_system = links + particles


# In[21]:


q = alpha + beta
u = omega + delta


# In[22]:


kane = me.KanesMethod(I, q_ind=q, u_ind=u, kd_eqs=kinematic_differentials)
fr, frstar = kane.kanes_equations(total_system, loads=forces)


# In[23]:


from numpy import radians, linspace, hstack, zeros, ones
from scipy.integrate import odeint
from pydy.codegen.ode_function_generators import generate_ode_function

param_syms = []
for par_seq in [l, m_bob, m_link, Ixx, Iyy, Izz, (g,)]:
   param_syms += list(par_seq)


# In[24]:


link_length = 10.0  # meters
link_mass = 10.0  # kg
link_radius = 0.5  # meters
link_ixx = 1.0 / 12.0 * link_mass * (3.0 * link_radius**2 + link_length**2)
link_iyy = link_mass * link_radius**2
link_izz = link_ixx

particle_mass = 5.0  # kg
particle_radius = 1.0  # meters


# In[25]:


param_vals = ([link_length for x in l] +
              [particle_mass for x in m_bob] +
              [link_mass for x in m_link] +
              [link_ixx for x in list(Ixx)] +
              [link_iyy for x in list(Iyy)] +
              [link_izz for x in list(Izz)] +
              [9.8])


# In[26]:


right_hand_side = generate_ode_function(kane.forcing_full, q, u, param_syms,
                                        mass_matrix=kane.mass_matrix_full,
                                        generator='cython',
                                        linear_sys_solver='sympy')


# In[27]:


duration = 10.0
fps = 60.0
t = linspace(0.0, duration, num=int(duration*fps))
x0 = hstack((ones(6) * radians(10.0), zeros(6)))

state_trajectories = odeint(right_hand_side, x0, t,
                            args=(dict(zip(param_syms, param_vals)),))


# In[28]:


from pydy.viz.shapes import Cylinder, Sphere
from pydy.viz.scene import Scene
from pydy.viz.visualization_frame import VisualizationFrame


# In[29]:


viz_frames = []

for i, (link, particle) in enumerate(zip(links, particles)):

   link_shape = Cylinder(name='cylinder{}'.format(i),
                        radius=link_radius,
                        length=link_length,
                        color='red')

   viz_frames.append(VisualizationFrame('link_frame{}'.format(i), link,
                                       link_shape))

   particle_shape = Sphere(name='sphere{}'.format(i),
                           radius=particle_radius,
                           color='blue')

   viz_frames.append(VisualizationFrame('particle_frame{}'.format(i),
                                       link.frame,
                                       particle,
                                       particle_shape))


# In[30]:


scene = Scene(I, O, *viz_frames)


# In[31]:


scene.times = t
scene.constants = dict(zip(param_syms, param_vals))
scene.states_symbols = q + u
scene.states_trajectories = state_trajectories

scene.display_jupyter()

