# Frames of the gear animation with Hertz contact from gears_hertz.jl. The left panel follows
# the contact point of the marked pair, 0.6 mm across, where the Hertz zone and the
# von Mises maximum below the surface show; the top right panel is the mesh zone of
# render.py with the window marked; the bottom right plot is the load share of a pair
# from the tooth stiffness against the 50/50 split. The contact window has a linear colour
# scale up to 500 MPa, which shows the subsurface maximum of about 0.56 p_max; the mesh zone
# has a logarithmic one, since the roots (about 130 MPa) and the contacts (400 to 600 MPa)
# are both readable only on the logarithm of the stress. The stress is the von Mises stress
# with σzz = ν(σxx + σyy), the state in the middle of the face width.
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
from matplotlib.colors import LogNorm, Normalize
import numpy as np

here = Path(__file__).parent
out = here/'out'
graphics = here.parent/'graphics'
fps = 24
BLUE = '#4472c4'

N1, N2, m = 20, 30, 4.0
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']
phiw = np.radians(summary['working_pressure_angle_deg'])
rw1 = rp1*np.cos(np.radians(20))/np.cos(phiw)
P = np.array([0.0, rw1])
d = np.array([-np.cos(phiw), np.sin(phiw)])
fr = np.genfromtxt(here/'frames2.csv', delimiter=',', names=True)
nframes = len(fr)
lognorm = LogNorm(vmin=10, vmax=1000)
vlin = 500


def load(name):
    X = np.loadtxt(out/f'{name}_hz_nodes.csv', delimiter=',')
    tris = np.loadtxt(out/f'{name}_hz_tris.csv', delimiter=',', dtype=int)
    vm = np.memmap(out/f'{name}_hz_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')


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]]


def draw_gear(ax, X, tris, vm, crop, log=True):
    x0, x1, y0, y1 = crop                  # only the triangles that show, for speed
    mx, my = 0.1*(x1 - x0), 0.1*(y1 - y0)
    inside = (X[:, 0] > x0 - mx) & (X[:, 0] < x1 + mx) & (X[:, 1] > y0 - my) & (X[:, 1] < y1 + my)
    tris = tris[inside[tris].any(axis=1)]
    if len(tris) == 0:                     # the gear has left the window
        return None
    tri = mtri.Triangulation(X[:, 0], X[:, 1], tris)
    if log:
        return ax.tripcolor(tri, np.clip(vm, 10, 1000), shading='gouraud', cmap='jet',
                            norm=lognorm, rasterized=True)
    return ax.tripcolor(tri, np.minimum(vm, vlin), shading='gouraud', cmap='jet',
                        vmin=0, vmax=vlin, rasterized=True)


def peak(W, vm, c, r):
    """World position and value of the largest stress within r of the point c."""
    near = np.flatnonzero(np.hypot(W[:, 0] - c[0], W[:, 1] - c[1]) < r)
    i = near[np.argmax(vm[near])]
    return W[i], vm[i]


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, 2)) for t in [lim[0], *inner, lim[1]]])


zoom = (-13.0, 13.0, 29.5, 51.5)           # mesh zone around the pitch point, as in render.py
hw, hh = 0.30, 0.36                        # half-width and half-height of the contact window


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))
    # The window follows the contact of the marked pair; before it touches, the pair ahead
    # of it, and after it lets go, the pair behind. Every pair repeats the marked one a base
    # pitch (40 frames) earlier or later, so its data are those of frame i ± 40.
    s0 = fr['s0'][i]
    j = 0 if s_a <= s0 <= s_r else (1 if s0 < s_a else -1)
    k = (i + 40*j) % nframes                # the frame at which the marked pair is here
    sf = s0 + j*pb
    c = P + sf*d
    win = (c[0] - hw, c[0] + hw, c[1] - hh, c[1] + hh)
    fig = plt.figure(figsize=(12.8, 7.2), facecolor='white')

    # Contact window
    ax = fig.add_axes([0.0, 0.0, 0.45, 1.0])
    draw_gear(ax, W1, T1, VM1[i], win, False); draw_gear(ax, W2, T2, VM2[i], win, False)
    im = plt.cm.ScalarMappable(norm=Normalize(0, vlin), cmap='jet')
    a = fr['a_mm'][k]
    for W, vm in ((W1, VM1[i]), (W2, VM2[i])):
        x, v = peak(W, vm, c, 4*a)
        ax.plot(*x, 'o', ms=7, mfc='none', mec='white', mew=1.6, zorder=5)
    who = {0: 'the marked pair', 1: 'the pair ahead of the marked one',
           -1: 'the pair behind the marked one'}[j]
    txt = (f"{who}: load share $f/F = {fr['share'][k]:.2f}$\n"
           f"Hertz: $a = {a*1e3:.0f}$ µm, $p_{{max}} = {fr['p_hertz_MPa'][k]:.0f}$ MPa\n"
           f"largest von Mises {fr['p_vm_contact'][k]:.0f} MPa, "
           f"{fr['p_vm_depth_mm'][k]*1e3:.0f} µm deep (pinion)")
    ax.text(0.03, 0.97, txt, transform=ax.transAxes, va='top', fontsize=10,
            bbox=dict(fc='white', ec='none', alpha=0.85))
    # scale bar of 0.1 mm
    x0, y0 = win[0] + 0.04, win[2] + 0.04
    ax.plot([x0, x0 + 0.1], [y0, y0], color='white', lw=3, solid_capstyle='butt')
    ax.text(x0 + 0.05, y0 + 0.015, '0.1 mm', color='white', ha='center', va='bottom', fontsize=10)
    ax.set_xlim(win[:2]); ax.set_ylim(win[2:]); ax.set_aspect('equal'); ax.axis('off')
    cax = fig.add_axes([0.455, 0.12, 0.012, 0.76])
    cb = fig.colorbar(im, cax=cax, extend='max')
    cb.set_ticks(np.linspace(0, vlin, 6))
    cb.set_label('von Mises stress (MPa)')

    # Mesh zone with the window marked
    ax = fig.add_axes([0.53, 0.36, 0.36, 0.62])
    draw_gear(ax, W1, T1, VM1[i], zoom); im = draw_gear(ax, W2, T2, VM2[i], zoom)
    ax.add_patch(plt.Rectangle((win[0], win[2]), 2*hw, 2*hh, fill=False, ec='k', lw=1.5, zorder=6))
    ax.add_patch(plt.Circle(c, 1.0, fill=False, ec='k', lw=0.8, zorder=6))
    for tip in (np.array([np.cos(psi1), np.sin(psi1)])*(rp1 + m),
                np.array([0, C]) + np.array([np.cos(psi2), np.sin(psi2)])*(rp2 + m)):
        ax.plot(*tip, 'o', ms=5, mfc='white', mec='k', mew=1.2, zorder=5)   # the marked teeth
    ax.set_xlim(zoom[:2]); ax.set_ylim(zoom[2:]); ax.set_aspect('equal'); ax.axis('off')

    cax = fig.add_axes([0.905, 0.40, 0.012, 0.54])
    cb = fig.colorbar(im, cax=cax, extend='both')
    ticks = [10, 30, 100, 300, 1000]
    cb.set_ticks(ticks); cb.set_ticklabels([str(t) for t in ticks]); cb.minorticks_off()
    cb.set_label('von Mises stress (MPa), log scale')

    # Load share of a pair along its contact
    ax = fig.add_axes([0.585, 0.085, 0.385, 0.22])
    s = fr['s0']
    for lo, hi in ((s_a, s_r - pb), (s_a + pb, s_r)):
        ax.axvspan(lo, hi, color='0.9', lw=0)
    on = fr['share'] > 0
    idx = np.flatnonzero(on)
    cut = np.flatnonzero((np.diff(idx) > 1) | (np.diff(fr['npairs'][idx]) != 0)) + 1
    for n, run in enumerate(np.split(idx, cut)):
        ax.plot(s[run], fr['share'][run], color=BLUE, lw=1.8,
                label="from the tooth stiffness" if n == 0 else None)
    ax.plot([s[0], s_a, s_a, s_r - pb, s_r - pb, s_a + pb, s_a + pb, s_r, s_r, s[-1]],
            [0, 0, 0.5, 0.5, 1, 1, 0.5, 0.5, 0, 0], color='k', ls='--', lw=1, label='50/50 split')
    ax.plot(sf, fr['share'][k], 'o', color=BLUE, ms=7, zorder=5)   # the pair in the window
    ax.set_xlim(s[0], s[-1]); ax.set_ylim(0, 1.15)
    end_ticks(ax)
    ax.set_yticks([0, 0.5, 1, 1.15]); ax.set_yticklabels(['0', '0.5', '1', '1.15'])
    ax.set_xlabel('contact position $s$ on the line of action (mm)', fontsize=9)
    ax.set_ylabel('load share $f/F$', fontsize=9)
    ax.tick_params(labelsize=8)
    ax.legend(fontsize=8, loc='upper left', frameon=False)
    ax.text(0.99, 0.97, f"mesh stiffness {fr['k_mesh_N_per_mm_um'][i]:.1f} N/(mm·µm)",
            transform=ax.transAxes, ha='right', va='top', fontsize=8)
    ax.set_title('load share of a tooth pair, dot: the pair in the window', fontsize=9)

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


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


if __name__ == '__main__':
    (out/'frames_hz').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()
    subprocess.run(['ffmpeg', '-y', '-loglevel', 'error', '-framerate', str(fps),
                    '-i', str(out/'frames_hz'/'f%04d.png'), '-c:v', 'libx264', '-pix_fmt', 'yuv420p',
                    '-crf', '18', '-movflags', '+faststart',
                    str(graphics/'GridapGears_hertz.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_hertz.png', dpi=150)
    subprocess.run(['ffmpeg', '-y', '-loglevel', 'error', '-i', str(graphics/'GridapGears_hertz.png'),
                    '-vf', 'scale=1280:-2', '-q:v', '3', str(graphics/'GridapGears_hertz_poster.jpg')],
                   check=True)
