#include <metal_stdlib>
using namespace metal;

struct Uniforms {
    float2 resolution;
    float  time;
    float  timeDelta;
    float2 mouse;
    uint   frame;
    uint   pad;
};

struct Params {
    float speed;
    float iterations;
    float zoom;
    float wander;
    float glow;
    float bands;
    float heat;
    float volume;
};

struct RKVertexOut {
    float4 position [[position]];
    float2 uv;
};

vertex RKVertexOut rk_vertex(uint vid [[vertex_id]]) {
    float2 p = float2((vid << 1) & 2, vid & 2);
    RKVertexOut o;
    o.position = float4(p * 2.0 - 1.0, 0.0, 1.0);
    o.uv = p;
    return o;
}

#line 1 "shader.metal"
// Ray Kiln · Episode 7 · Fractal Glow
//
// @episode 7 Fractal Glow
// @length 60
// @variation 30
// @variations 1,2,4,6,7,8,12
// @still 40.0
// @teaches Complex numbers, Julia sets, smooth iteration colour
// @category Fragment
// @tags complex numbers, julia sets, fractals, smooth iteration count, distance estimation, procedural sound
// @final Fractal glow | Colour by the smooth count, let the edge of the set glow, and light the inside with the orbits.
//
// A Julia set: every pixel is a complex number, squared and added to over and over, and coloured by how quickly it
// flies away. The number added, c, walks slowly round the edge of the Mandelbrot set, so the shape keeps changing.
// Each @param line below becomes a slider in the host app and a row in the page's parameters table:
//   // @param name default min max description

// @param speed 0.5 0.0 2.0 How fast c walks round its path and the picture turns
// @param iterations 96.0 16.0 256.0 How many times each point is squared before it counts as staying
// @param zoom 0.75 0.4 3.0 How much of the plane fits on screen: higher zooms out
// @param wander 0.5 0.0 1.0 How far c strays inside and outside the edge, from whole shapes to scattered dust
// @param glow 0.5 0.0 1.0 How brightly the edge of the set glows
// @param bands 0.0 0.0 1.0 Brings back the steps of the plain iteration count
// @param heat 0.5 0.0 1.0 Slides the palette from deep ember to white-hot
// @param volume 0.8 0.0 1.0 Loudness of the sound

constant float TAU = 6.2831853;

// ---- Complex numbers -------------------------------------------------------------------------------------------
// A complex number is a pair (x, y), written x + iy, where i * i = -1. Adding works on each half; multiplying turns
// and stretches: the angles add and the lengths multiply. Squaring a number doubles its angle and squares its length.

float2 cmul(float2 a, float2 b) {
    return float2(a.x * b.x - a.y * b.y, a.x * b.y + a.y * b.x);
}

float2 csqr(float2 z) {
    return float2(z.x * z.x - z.y * z.y, 2.0 * z.x * z.y);
}

// Episode 3's rotation.
float2 rot(float2 p, float a) {
    float c = cos(a), s = sin(a);
    return float2(c * p.x - s * p.y, s * p.x + c * p.y);
}

// ---- The path of c ---------------------------------------------------------------------------------------------
// The Mandelbrot set is every c whose Julia set holds together. Its biggest part is a heart shape, the cardioid, and
// the edge of it is w / 2 - w * w / 4 for w going once round the unit circle. Julia sets for c on that edge are the
// most intricate; a little inside they fill in, a little outside they break into dust. The sound uses this too.

float2 juliaC(float t, float wander) {
    float a = t * 0.25 + 2.2;
    float2 w = float2(cos(a), sin(a));
    float2 c = 0.5 * w - 0.25 * csqr(w);
    return c * (1.0 + wander * 0.07 * sin(t * 0.43));
}

// ---- The kiln palette from episode 1 ---------------------------------------------------------------------------

float3 kiln(float t) {
    float3 c0 = float3(0.020, 0.012, 0.020);
    float3 c1 = float3(0.280, 0.040, 0.030);
    float3 c2 = float3(0.880, 0.260, 0.050);
    float3 c3 = float3(1.000, 0.680, 0.200);
    float3 c4 = float3(1.000, 0.970, 0.840);
    t = clamp(t, 0.0, 1.0);
    float3 c = mix(c0, c1, smoothstep(0.00, 0.22, t));
    c = mix(c, c2, smoothstep(0.22, 0.50, t));
    c = mix(c, c3, smoothstep(0.50, 0.76, t));
    return mix(c, c4, smoothstep(0.76, 1.00, t));
}

float4 shade(float2 uv, constant Uniforms& u, constant Params& p) {
    float t = u.time * p.speed;
    float px = 3.0 * p.zoom / min(u.resolution.x, u.resolution.y);
    float2 z = rot((uv - 0.5) * u.resolution * px, t * 0.08);
    float2 c = juliaC(t, p.wander);
    int n = int(clamp(p.iterations, 1.0, 256.0));

    // The loop also keeps dz, how fast z changes as the starting point moves. Each step multiplies it by 2 * z.
    float2 dz = float2(1.0, 0.0);
    // And it keeps trap, how close the orbit comes to the centre.
    float trap = 1e9;
    float steps = 0.0;
    bool escaped = false;
    for (int i = 0; i < n; i++) {
        dz = 2.0 * cmul(z, dz);
        z = csqr(z) + c;
        trap = min(trap, length(z));
        if (dot(z, z) > 256.0 * 256.0) { escaped = true; break; }
        steps += 1.0;
    }

    if (!escaped) {
        // Inside the set, the nearer the orbit came to the centre, the hotter the embers.
        float ember = exp(-trap * 4.0);
        float level = 0.06 + 0.42 * ember;
        return float4(kiln(level + (p.heat - 0.5) * 0.5), 1.0);
    }

    // The smooth count: log2(log2|z|) grows by one with each step once |z| is large, so taking it away from the
    // count joins the steps up into a smooth slope.
    float r = length(z);
    float smoothN = steps + 1.0 - log2(log2(r));
    float stepped = mix(smoothN, floor(smoothN), p.bands);
    float level = 0.9 * pow(clamp(stepped / 40.0, 0.0, 1.0), 0.6);

    // The distance estimate: how far this pixel is from the set, from |z| and how fast it grew. Within a few pixels
    // of the edge, it glows.
    float dist = 0.5 * r * log(r) / max(length(dz), 1e-20);
    float edge = exp(-dist / (px * (1.0 + 4.0 * p.glow)));
    level = max(level, edge * (0.2 + 0.8 * p.glow));

    return float4(kiln(level + (p.heat - 0.5) * 0.5), 1.0);
}

// ---- Sound -----------------------------------------------------------------------------------------------------
// The sound follows the orbit of zero, the point every Julia set is built around: start at 0 and keep squaring and
// adding c, as the picture does for every pixel. Each of the first eight steps is a note of A minor pentatonic, picked
// by the angle of the point and heard from where it lies left or right, played in turn as an arpeggio. When c strays
// outside the set, the orbit flies off and its later notes fall silent, so the music thins as the picture turns to
// dust. Under it a low A hums, brighter as more of the orbit stays.

// Episode 4's plucked note: a quick rise, then a fall, with a little of its octave.
float pluck(float f, float age) {
    float env = (1.0 - exp(-age * 300.0)) * exp(-age * 9.0);
    return (sin(TAU * f * age) + 0.3 * sin(2.0 * TAU * f * age)) * env;
}

// How much of the orbit of zero stays near, from 0 to 1: all of it while c is in the set.
float orbitStay(float t, constant Params& p) {
    float2 c = juliaC(t * p.speed, p.wander);
    float2 z = float2(0.0);
    float stay = 0.0;
    for (int i = 0; i < 16; i++) {
        z = csqr(z) + c;
        if (dot(z, z) > 4.0) break;
        stay += 1.0 / 16.0;
    }
    return stay;
}

float2 sound(float t, constant Params& p) {
    const float scale[10] = { 220.00, 261.63, 293.66, 329.63, 392.00, 440.00, 523.25, 587.33, 659.25, 783.99 };
    const float rate = 4.0;

    float2 notes = float2(0.0);
    // The note sounding now and the three before it, still fading.
    float now = floor(t * rate);
    for (int back = 0; back < 4; back++) {
        float idx = now - float(back);
        if (idx < 0.0) continue;
        float start = idx / rate;
        float age = t - start;
        int k = int(fmod(idx, 8.0));
        // The orbit of zero for c as it was when the note started, k + 1 steps in.
        float2 c = juliaC(start * p.speed, p.wander);
        float2 z = float2(0.0);
        float alive = 1.0;
        for (int i = 0; i <= k; i++) {
            z = csqr(z) + c;
            if (dot(z, z) > 4.0) { alive = 0.0; break; }
        }
        float a = atan2(z.y, z.x) / TAU + 0.5;
        float f = scale[int(clamp(floor(a * 10.0), 0.0, 9.0))];
        float pan = 0.5 + 0.4 * clamp(z.x, -1.0, 1.0);
        float s = pluck(f, age) * alive * 0.22;
        notes += float2(1.0 - pan, pan) * s;
    }

    // How much of the orbit of zero stays near, measured four times a second and blended between, so it never jumps.
    float q = t * 4.0;
    float stay = mix(orbitStay(floor(q) / 4.0, p), orbitStay(floor(q) / 4.0 + 0.25, p), fract(q));
    float b = 0.15 + 0.6 * stay;
    float2 low = float2(sin(TAU * 55.0 * 1.002 * t) + b * 0.5 * sin(TAU * 110.0 * 1.002 * t) + b * 0.25 * sin(TAU * 165.0 * t),
                        sin(TAU * 55.0 * 0.998 * t) + b * 0.5 * sin(TAU * 110.0 * 0.998 * t + 0.3) + b * 0.25 * sin(TAU * 165.0 * t + 0.6))
                 * 0.12;

    float fadeIn = smoothstep(0.0, 2.0, t);
    return tanh((notes + low) * 2.0) * 0.58 * fadeIn * p.volume;
}

#line 1 "rk_host"
fragment float4 rk_fragment(RKVertexOut in [[stage_in]],
                                constant Uniforms& u [[buffer(0)]],
                                constant Params& p [[buffer(1)]]) {
    return shade(in.uv, u, p);
}

struct SoundUniforms { float start; float rate; uint count; uint pad; };
kernel void rk_sound(device float2* out [[buffer(0)]],
                         constant SoundUniforms& su [[buffer(1)]],
                         constant Params& p [[buffer(2)]],
                         uint i [[thread_position_in_grid]]) {
    if (i >= su.count) return;
    out[i] = clamp(sound(su.start + float(i) / su.rate, p), -1.0, 1.0);
}