# Frames of the gear stress animation from the gears.jl results: the mesh zone in von Mises
# stress, the whole pair beside it so the rotation shows, and the root stress of the tracked
# tooth pair against the contact position, with a marker that follows the frame. Then
# ffmpeg turns the frames into the looping video, and one single-contact frame is the poster.
import subprocess
import sys
from multiprocessing import Pool
from pathlib import Path

import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import matplotlib.tri as mtri
import numpy as np

here = Path(__file__).parent
out = here/'out'
graphics = here.parent/'graphics'
fps = 24

N1, N2, m, phi = 20, 30, 4.0, np.radians(20)
rp1, rp2 = m*N1/2, m*N2/2
summary = dict(np.genfromtxt(here/'summary.csv', delimiter=',', dtype=None, encoding=None))
C = summary['centre_distance_mm']
s_a, s_r, pb = summary['s_a_mm'], summary['s_r_mm'], summary['base_pitch_mm']
fr = np.genfromtxt(out/'frames.csv', delimiter=',', names=True)
nframes = len(fr)

def load(name):
    X = np.loadtxt(out/f'{name}_nodes.csv', delimiter=',')
    tris = np.loadtxt(out/f'{name}_tris.csv', delimiter=',', dtype=int)
    vm = np.memmap(out/f'{name}_vm.f32', dtype=np.float32, mode='r').reshape(nframes, len(X))
    return X, tris, vm

X1, T1, VM1 = load('pinion')
X2, T2, VM2 = load('wheel')

# A fixed colour range for all frames, clipped just above the largest fillet stress so the
# roots are resolved; the contact peaks saturate
vmax = 10*np.ceil(1.05*max(fr['p_vm_comp'].max(), fr['w_vm_comp'].max())/10)

def place(X, psi, centre):
    """Body-frame nodes turned by psi and moved to the gear centre (world frame)."""
    c, s = np.cos(psi), np.sin(psi)
    return np.c_[c*X[:, 0] - s*X[:, 1] + centre[0], s*X[:, 0] + c*X[:, 1] + centre[1]]

zoom = (-13.0, 13.0, 29.5, 51.5)            # mesh zone around the pitch point (0, rp1)

def draw_gear(ax, X, tris, vm, crop=None):
    if crop is not None:                     # only the triangles that show, for speed
        x0, x1, y0, y1 = crop
        inside = ((X[:, 0] > x0 - 2) & (X[:, 0] < x1 + 2) & (X[:, 1] > y0 - 2) & (X[:, 1] < y1 + 2))
        tris = tris[inside[tris].any(axis=1)]
    tri = mtri.Triangulation(X[:, 0], X[:, 1], tris)
    return ax.tripcolor(tri, np.minimum(vm, vmax), shading='gouraud', cmap='jet',
                        vmin=0, vmax=vmax, rasterized=True)

def tooth_tip(psi, centre, ra):
    """World position of the tip of tooth 0 (the tracked one) of a gear."""
    return np.array(centre) + ra*np.array([np.cos(psi), np.sin(psi)])

def end_ticks(ax, fmt='{:g}'):
    """Label the ends of both axes as well as the inner ticks."""
    for axis, lim in ((ax.xaxis, ax.get_xlim()), (ax.yaxis, ax.get_ylim())):
        ticks = [t for t in axis.get_ticklocs() if lim[0] < t < lim[1]]
        inner = [t for t in ticks if abs(t - lim[0]) > 0.1*(lim[1] - lim[0])
                 and abs(t - lim[1]) > 0.1*(lim[1] - lim[0])]
        axis.set_ticks([lim[0], *inner, lim[1]])
        axis.set_ticklabels([fmt.format(round(t, 1)) for t in [lim[0], *inner, lim[1]]])

def frame(i, path, dpi=100):
    psi1, psi2 = fr['psi1'][i], fr['psi2'][i]
    W1, W2 = place(X1, psi1, (0, 0)), place(X2, psi2, (0, C))
    fig = plt.figure(figsize=(12.8, 7.2), facecolor='white')

    # Mesh zone
    ax = fig.add_axes([0.0, 0.0, 0.56, 1.0])
    draw_gear(ax, W1, T1, VM1[i], zoom); im = draw_gear(ax, W2, T2, VM2[i], zoom)
    for tip in (tooth_tip(psi1, (0, 0), rp1 + m), tooth_tip(psi2, (0, C), rp2 + m)):
        ax.plot(*tip, 'o', ms=5, mfc='white', mec='k', mew=1.2, zorder=5)
    ax.set_xlim(zoom[:2]); ax.set_ylim(zoom[2:]); ax.set_aspect('equal'); ax.axis('off')

    cax = fig.add_axes([0.575, 0.12, 0.012, 0.76])
    cb = fig.colorbar(im, cax=cax, extend='max')
    cb.set_ticks(np.linspace(0, vmax, 5))
    cb.set_ticklabels([f'{t:.0f}' for t in np.linspace(0, vmax, 5)])
    cb.set_label('von Mises stress (MPa)')

    # Whole pair, with the zoom window and the line of action
    ax = fig.add_axes([0.655, 0.40, 0.33, 0.58])
    draw_gear(ax, W1, T1, VM1[i]); draw_gear(ax, W2, T2, VM2[i])
    ax.add_patch(plt.Rectangle((zoom[0], zoom[2]), zoom[1] - zoom[0], zoom[3] - zoom[2],
                               fill=False, ec='k', lw=1.2))
    for ctr in ((0, 0), (0, C)):
        ax.plot(*ctr, '+', color='k', ms=6)
    ax.set_xlim(-70, 70); ax.set_ylim(-60, 175); ax.set_aspect('equal'); ax.axis('off')
    ax.text(0, 168, 'wheel, $z_2 = 30$', ha='center', fontsize=9)
    ax.text(0, -58, 'pinion, $z_1 = 20$, drives with $T_1 = 100$ N$\\cdot$m', ha='center', fontsize=9)

    # Root stress of the tracked pair (white dots) against its contact position
    ax = fig.add_axes([0.665, 0.085, 0.31, 0.27])
    s = fr['s0']
    for lo, hi, shade in [(s_a, s_r - pb, '0.88'), (s_r - pb, s_a + pb, '0.96'), (s_a + pb, s_r, '0.88')]:
        ax.axvspan(lo, hi, color=shade, lw=0)
    ax.plot(s, fr['p_sig1_tens'], color='C0', lw=1.5, label='pinion')
    ax.plot(s, fr['w_sig1_tens'], color='C3', lw=1.5, label='wheel')
    ax.axhline(summary['pinion_lewis_MPa'], color='C0', ls='--', lw=0.8)
    ax.axhline(summary['wheel_lewis_MPa'], color='C3', ls='--', lw=0.8)
    ax.plot(s[i], fr['p_sig1_tens'][i], 'o', color='C0', ms=6)
    ax.plot(s[i], fr['w_sig1_tens'][i], 'o', color='C3', ms=6)
    ax.set_xlim(s[0], s[-1]); ax.set_ylim(0, 10*np.ceil(1.1*fr['p_sig1_tens'].max()/10))
    end_ticks(ax)
    ax.set_xlabel('contact position $s$ on the line of action (mm)', fontsize=9)
    ax.set_ylabel('root stress $\\sigma_1$ (MPa)', fontsize=9)
    ax.tick_params(labelsize=8)
    ax.legend(fontsize=8, loc='upper left', frameon=False)
    ax.text(s[-1], summary['pinion_lewis_MPa'] + 2, 'Lewis ', ha='right', va='bottom', fontsize=8, color='C0')
    ax.set_title('root stress of the marked teeth (tension fillet)', fontsize=9)
    ax.text(0.5*(s_r - pb + s_a + pb), 3, 'single', ha='center', fontsize=8)
    for x in (0.5*(s_a + s_r - pb), 0.5*(s_a + pb + s_r)):
        ax.text(x, 3, 'double', ha='center', fontsize=8)

    fig.savefig(path, dpi=dpi, facecolor='white')
    plt.close(fig)

def render(i):
    frame(i, out/'frames'/f'f{i:04d}.png')

if __name__ == '__main__':
    (out/'frames').mkdir(exist_ok=True)
    only = [int(a) for a in sys.argv[1:]]
    with Pool(4) as pool:
        pool.map(render, only or range(nframes))
    if only:
        sys.exit()
    graphics.mkdir(exist_ok=True)
    subprocess.run(['ffmpeg', '-y', '-loglevel', 'error', '-framerate', str(fps),
                    '-i', str(out/'frames'/'f%04d.png'), '-c:v', 'libx264', '-pix_fmt', 'yuv420p',
                    '-crf', '18', '-movflags', '+faststart',
                    str(graphics/'GridapGears_stress.mp4')], check=True)
    # Poster: the single-contact frame closest to the pitch point
    single = np.flatnonzero(fr['npairs'] == 1)
    i = single[np.argmin(np.abs(fr['s0'][single]))]
    frame(i, graphics/'GridapGears_stress.png', dpi=150)
    subprocess.run(['ffmpeg', '-y', '-loglevel', 'error', '-i', str(graphics/'GridapGears_stress.png'),
                    '-vf', 'scale=1280:-2', '-q:v', '3', str(graphics/'GridapGears_stress_poster.jpg')],
                   check=True)
