#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 balls;
    float melt;
    float rise;
    float soft;
    float occlusion;
    float light;
    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 6 · Molten Metaballs
//
// @episode 6 Molten Metaballs
// @length 60
// @variation 30
// @variations 1,2,3,4,5,6,8,12
// @still 14.0
// @teaches 3D smooth blending, soft shadows, ambient occlusion
// @category Fragment
// @tags raymarching, signed distance fields, smooth minimum, metaballs, soft shadows, ambient occlusion, procedural sound
// @final Molten | Colour it with the kiln, let the necks where blobs melt together glow, and warm the floor beneath them.
//
// Blobs of molten metal drifting over a kiln floor, melting into each other and into the floor wherever they meet.
// Their shadows on the floor are soft at the edges, and the creases between them darken because little light gets in.
// 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 the blobs drift and the camera circles
// @param balls 5.0 2.0 7.0 How many blobs there are
// @param melt 0.5 0.0 1.0 How far apart blobs start to melt into each other and into the floor
// @param rise 1.0 0.0 1.5 How high the blobs float up and how deep they sink into the floor
// @param soft 0.5 0.0 1.0 How soft the shadows are at their edges: 0 is a hard edge
// @param occlusion 1.0 0.0 2.0 How dark the creases get where little light reaches
// @param light 0.15 0.0 1.0 Where the light stands, once around the scene
// @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

// ---- Shapes ----------------------------------------------------------------------------------------------------

float sdSphere(float3 p, float r) {
    return length(p) - r;
}

// Episode 2's smooth minimum: within k of each other, two distances are pulled down together.
float smin(float a, float b, float k) {
    float h = max(k - abs(a - b), 0.0) / k;
    return min(a, b) - h * h * k * 0.25;
}

// Where blob i is at time t. Each drifts around the middle on its own slow orbit, and bobs up and down.
// The sound uses these too.
float3 ballPos(int i, float t, float rise) {
    float fi = float(i);
    float a = t * (0.30 + 0.07 * fi) + fi * 2.4;
    float r = 0.9 + 0.45 * sin(t * (0.23 + 0.05 * fi) + fi * 2.1);
    float y = 0.15 + rise * 0.5 * sin(t * (0.55 + 0.11 * fi) + fi * 1.7);
    return float3(r * cos(a), y, r * sin(a));
}

float ballRadius(int i) {
    return 0.34 + 0.08 * sin(float(i) * 3.7);
}

constant float FLOOR = -0.55;

// The blobs melt into each other: fold the smooth minimum over them, one at a time.
float blobs(float3 p, float t, int n, float rise, float k) {
    float d = 1e9;
    for (int i = 0; i < n; i++) {
        d = smin(d, sdSphere(p - ballPos(i, t, rise), ballRadius(i)), k);
    }
    return d;
}

// The same blobs joined with a plain min, to measure how much the smooth minimum pulled them together.
float blobsHard(float3 p, float t, int n, float rise) {
    float d = 1e9;
    for (int i = 0; i < n; i++) {
        d = min(d, sdSphere(p - ballPos(i, t, rise), ballRadius(i)));
    }
    return d;
}

// Then the blobs melt into the floor too, a little more tightly, so they pool where they touch it.
float map(float3 p, float t, int n, float rise, float k) {
    return smin(blobs(p, t, n, rise, k), p.y - FLOOR, k * 0.6);
}

// The normal, as in episode 3: how the distance changes for a small step along each axis.
float3 normalAt(float3 p, float t, int n, float rise, float k) {
    const float2 e = float2(0.001, 0.0);
    return normalize(float3(
        map(p + e.xyy, t, n, rise, k) - map(p - e.xyy, t, n, rise, k),
        map(p + e.yxy, t, n, rise, k) - map(p - e.yxy, t, n, rise, k),
        map(p + e.yyx, t, n, rise, k) - map(p - e.yyx, t, n, rise, k)));
}

// ---- Light -----------------------------------------------------------------------------------------------------

// A soft shadow: march from the surface towards the light. The nearer the ray passes to something, compared with how
// far it has gone, the deeper the shadow, so a shadow fades out over a band instead of stopping at a line.
float softShadow(float3 ro, float3 rd, float w,
                 float t, int n, float rise, float k) {
    float res = 1.0;
    float s = 0.02;
    for (int i = 0; i < 48; i++) {
        float h = map(ro + rd * s, t, n, rise, k);
        res = min(res, h / (w * s));
        if (res < 0.001 || s > 6.0) break;
        s += clamp(h, 0.01, 0.3);
    }
    res = clamp(res, 0.0, 1.0);
    return res * res * (3.0 - 2.0 * res);
}

// Ambient occlusion: step out along the normal. Where the scene is nearer than the step, something blocks the light
// from that side, and the nearer steps count most.
float ambientOcclusion(float3 pos, float3 nrm,
                       float t, int n, float rise, float k) {
    float occ = 0.0;
    float weight = 1.0;
    for (int i = 0; i < 5; i++) {
        float h = 0.02 + 0.12 * float(i);
        occ += (h - map(pos + nrm * h, t, n, rise, k)) * weight;
        weight *= 0.8;
    }
    return clamp(1.0 - 2.0 * occ, 0.0, 1.0);
}

// The camera circles the scene slowly, above it, looking down at the middle of the floor.
float3 cameraPos(float t) {
    float a = t * 0.2;
    return float3(3.4 * sin(a), 1.6, 3.4 * cos(a));
}

// 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;
    int n = int(clamp(round(p.balls), 1.0, 7.0));
    float k = max(p.melt * 0.9, 0.0001);

    float side = min(u.resolution.x, u.resolution.y);
    float2 q = (uv - 0.5) * u.resolution / side * 2.0;
    float3 ro = cameraPos(t);
    float3 fw = normalize(float3(0.0, -0.1, 0.0) - ro);
    float3 rt = normalize(cross(fw, float3(0.0, 1.0, 0.0)));
    float3 up = cross(rt, fw);
    float3 rd = normalize(q.x * rt + q.y * up + 2.0 * fw);

    float dist = 0.0;
    bool hit = false;
    for (int i = 0; i < 128; i++) {
        float d = map(ro + rd * dist, t, n, p.rise, k);
        if (d < 0.001 * dist) { hit = true; break; }
        dist += d;
        if (dist > 20.0) break;
    }

    // Behind everything, the dark of the kiln.
    if (!hit) return float4(kiln(0.02 + (p.heat - 0.5) * 0.5), 1.0);

    float3 pos = ro + rd * dist;
    float3 nrm = normalAt(pos, t, n, p.rise, k);
    float a = p.light * 6.2831853;
    float3 l = normalize(float3(cos(a), 1.4, sin(a)));
    float diffuse = max(dot(nrm, l), 0.0);
    float w = 0.005 + 0.3 * p.soft;
    float shadow = softShadow(pos + nrm * 0.01, l, w, t, n, p.rise, k);
    float ao = ambientOcclusion(pos, nrm, t, n, p.rise, k);
    ao = mix(1.0, ao, clamp(p.occlusion, 0.0, 1.0));
    ao = pow(ao, max(p.occlusion, 1.0));
    // A highlight, as in episode 3.
    float3 h = normalize(l - rd);
    float spec = pow(max(dot(nrm, h), 0.0), 48.0) * diffuse * shadow;

    // On a blob or on the floor? Whichever surface is nearer.
    float db = blobs(pos, t, n, p.rise, k);
    float metal = smoothstep(0.05, -0.05, db - (pos.y - FLOOR));

    // How far smin pulled the surface out: a quarter of k mid-neck.
    float pull = (blobsHard(pos, t, n, p.rise) - db) / (k * 0.25);
    float neck = smoothstep(0.45, 1.0, pull);

    // The floor is dark clay, warmed by the metal above it.
    float bounce = exp(-max(db, 0.0) * 3.0) * 0.3;
    float floorLevel = (0.04 + 0.30 * diffuse * shadow + bounce) * ao;
    // The metal is lit, has a highlight, and glows where it melts.
    float metalLevel = (0.24 + 0.40 * diffuse * shadow) * ao;
    metalLevel += 0.5 * spec + 0.4 * neck;

    float level = mix(floorLevel, metalLevel, metal);
    // Far away the floor fades into the dark.
    level = mix(level, 0.02, smoothstep(3.5, 9.0, dist));
    float3 col = kiln(level + (p.heat - 0.5) * 0.5);
    return float4(col, 1.0);
}

// ---- Sound -----------------------------------------------------------------------------------------------------
// The sound is built from the blobs' positions, the same functions the picture uses. Under everything, a low E hums
// and grows brighter the more the blobs melt together. Each blob has a note of E minor pentatonic, heard where it is
// on the screen and only when it comes close enough to another blob to start melting into it, so chords swell as blobs
// meet. Above them the metal bubbles: short pops that rise in pitch, more of them the softer the melt.

// Integer hash, so the bubbles stay exact for minutes of sound.
float hashU(uint n) {
    n = (n << 13u) ^ n;
    n = n * (n * n * 15731u + 789221u) + 1376312589u;
    return float(n & 0x7fffffffu) / float(0x7fffffff);
}

// A soft tone with some of its overtones: `bright` from 0 (a pure sine) to 1 (a fuller sound).
float voice(float ph, float bright) {
    const float tau = 6.2831853;
    return sin(tau * ph) + bright * (0.45 * sin(2.0 * tau * ph) + 0.25 * sin(3.0 * tau * ph) + 0.12 * sin(4.0 * tau * ph));
}

float2 sound(float t, constant Params& p) {
    const float tau = 6.2831853;
    const float notes[7] = { 164.81, 196.00, 220.00, 246.94, 293.66, 329.63, 392.00 };
    float ts = t * p.speed;
    int n = int(clamp(round(p.balls), 1.0, 7.0));
    float k = max(p.melt * 0.9, 0.0001);

    float3 ro = cameraPos(ts);
    float3 fw = normalize(float3(0.0, -0.1, 0.0) - ro);
    float3 rt = normalize(cross(fw, float3(0.0, 1.0, 0.0)));

    // Each blob's note, as loud as it is close to melting into its nearest neighbour.
    float2 chord = float2(0.0);
    float meltAll = 0.0;
    for (int i = 0; i < n; i++) {
        float3 pi = ballPos(i, ts, p.rise);
        float gap = 10.0;
        for (int j = 0; j < n; j++) {
            if (j == i) continue;
            gap = min(gap, length(pi - ballPos(j, ts, p.rise)) - ballRadius(i) - ballRadius(j));
        }
        float close = exp(-max(gap, 0.0) * 3.0 / (0.15 + k));
        meltAll += close;
        float3 rel = pi - ro;
        float x = clamp(dot(rel, rt) / max(dot(rel, fw), 0.5) * 0.9, -1.0, 1.0);
        float pan = 0.5 + 0.4 * x;
        float tone = voice(notes[i] * t, 0.15 + 0.5 * close) * close * 0.05;
        chord += float2(1.0 - pan, pan) * tone * 1.4;
    }
    meltAll /= float(n);

    // The low E, a few cents apart in each ear, brighter as the blobs melt.
    float b = 0.1 + 0.8 * meltAll;
    float2 low = float2(voice(41.20 * 1.002 * t, b) + 0.6 * voice(82.41 * 0.999 * t, b),
                        voice(41.20 * 0.998 * t, b) + 0.6 * voice(82.41 * 1.001 * t + 0.25, b)) * 0.12;

    // Bubbles: time cut into slots; some slots get a pop at a random moment, rising from a random pitch.
    float2 bubbles = float2(0.0);
    const float slot = 0.17;
    float chance = 0.25 + 0.35 * p.melt;
    for (int back = 0; back < 2; back++) {
        float idx = floor(t / slot) - float(back);
        if (idx < 0.0) continue;
        uint s = uint(idx);
        if (hashU(s * 3u) > chance) continue;
        float tau0 = t - (idx + 0.6 * hashU(s * 3u + 1u)) * slot;
        if (tau0 < 0.0) continue;
        float f0 = 180.0 + 420.0 * hashU(s * 3u + 2u);
        float env = (1.0 - exp(-tau0 * 600.0)) * exp(-tau0 * 28.0);
        float pop = sin(tau * f0 * (tau0 + 5.0 * tau0 * tau0)) * env;
        float pan = 0.2 + 0.6 * hashU(s * 7u + 5u);
        bubbles += float2(1.0 - pan, pan) * pop * 0.22;
    }

    float fadeIn = smoothstep(0.0, 2.0, t);
    return tanh((low + chord + bubbles) * 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);
}