#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 lift;
    float blend;
    float light;
    float shine;
    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 3 · Floating 3D Scene
//
// @episode 3 Floating 3D Scene
// @length 60
// @variation 30
// @variations 2,3,6,7,9,10,11,12
// @still 21.0
// @teaches Raymarching, normals, basic lighting
// @category Fragment
// @tags raymarching, signed distance fields, normals, lighting, procedural sound
// @final Kiln light | A warm light from above, the kiln's glow from below, a highlight, and a camera that circles the scene.
//
// A sphere, a ring and a tumbling box floating in the dark, drawn by marching a ray from the camera through every
// pixel until it meets a surface, then lit by how that surface faces the light.
// 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 shapes float and the camera circles
// @param lift 1.0 0.0 2.0 How far the shapes bob up and down
// @param blend 0.3 0.0 1.0 How softly the ring melts into the sphere
// @param light 0.15 0.0 1.0 Where the warm light stands, once around the scene
// @param shine 40.0 4.0 160.0 How small and sharp the highlights are
// @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 ----------------------------------------------------------------------------------------------------
// The same idea as episode 2, one dimension up: each shape is a function that says how far a point is from it.

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

// A ring lying flat: R is its radius, r the thickness of its tube.
float sdTorus(float3 p, float R, float r) {
    float2 q = float2(length(p.xz) - R, p.y);
    return length(q) - r;
}

float sdBox(float3 p, float3 b) {
    float3 d = abs(p) - b;
    return length(max(d, 0.0)) + min(max(d.x, max(d.y, d.z)), 0.0);
}

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;
}

// Turn a pair of coordinates by angle a.
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);
}

// Where each shape is at time t. The sound uses these too.
float3 spherePos(float t, float lift) {
    return float3(0.0, 0.22 * lift * sin(t * 1.3), 0.0);
}

float3 ringPos(float t, float lift) {
    return float3(0.0, 0.22 * lift * sin(t * 1.3 - 1.4), 0.0);
}

float3 boxPos(float t, float lift) {
    float a = t * 0.5;
    float y = 0.3 + 0.3 * lift * sin(t * 1.7);
    return float3(1.9 * cos(a), y, 1.9 * sin(a));
}

// The whole scene: the distance from p to the nearest surface.
float map(float3 p, float t, float lift, float k) {
    float d = sdSphere(p - spherePos(t, lift), 0.6);
    // The ring tilts back and forth and spins around its own axis.
    float3 q = p - ringPos(t, lift);
    q.yz = rot(q.yz, 0.45 + 0.35 * sin(t * 0.7));
    q.xz = rot(q.xz, t * 0.6);
    d = smin(d, sdTorus(q, 0.86, 0.12), k);
    // The box orbits the middle and tumbles as it goes, its corners rounded off.
    float3 b = p - boxPos(t, lift);
    b.xy = rot(b.xy, t * 0.9);
    b.yz = rot(b.yz, t * 0.7);
    d = min(d, sdBox(b, float3(0.26)) - 0.05);
    return d;
}

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

// The camera circles the scene slowly, a little above it.
float3 cameraPos(float t) {
    float a = t * 0.2;
    return float3(3.9 * sin(a), 1.2, 3.9 * 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;
    float k = max(p.blend * 0.5, 0.0001);

    // A ray for this pixel, from a camera that looks at the middle of the scene.
    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(-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);

    // March: step as far as the nearest surface, again and again, until the ray touches one or leaves the scene.
    // Remember how close the ray came to anything, for the glow around the shapes.
    float dist = 0.0;
    float closest = 1e9;
    bool hit = false;
    for (int i = 0; i < 128; i++) {
        float d = map(ro + rd * dist, t, p.lift, k);
        closest = min(closest, d);
        if (d < 0.001 * dist) { hit = true; break; }
        dist += d;
        if (dist > 20.0) break;
    }

    // Behind the shapes: dark, warmer towards the bottom where the kiln is hottest, with a glow around each shape.
    float low = smoothstep(0.1, -0.8, rd.y);
    float halo = exp(-max(closest, 0.0) * 10.0);
    float level = 0.03 + 0.12 * low + 0.28 * halo;

    if (hit) {
        float3 pos = ro + rd * dist;
        float3 n = normalAt(pos, t, p.lift, k);
        // The warm light, high up, on a circle around the scene.
        float a = p.light * 6.2831853;
        float3 l = normalize(float3(cos(a), 1.1, sin(a)));
        // Diffuse: a surface is as bright as it faces the light.
        float diffuse = max(dot(n, l), 0.0);
        // The kiln's glow, rising from below onto every surface that faces down.
        float below = max(-n.y, 0.0);
        // A highlight where the surface would mirror the light into the camera.
        float3 h = normalize(l - rd);
        float spec = pow(max(dot(n, h), 0.0), p.shine) * step(0.0, dot(n, l));
        level = 0.12 + 0.55 * diffuse + 0.22 * below + 0.6 * spec;
    }

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

// ---- Sound -----------------------------------------------------------------------------------------------------
// Each shape has a voice, and the voices follow the shapes. The sphere hums a low D, the ring an A above it, and the
// box an F-sharp higher still: a D major chord. A voice grows brighter as its shape floats up towards the light, the
// box is heard from where it is on the screen and louder as it swings close to the camera, and the ring shimmers as
// it melts into the sphere.

// A soft tone with some of its overtones: `bright` from 0 (a pure sine) to 1 (a reedy, 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;
    float ts = t * p.speed;
    float k = max(p.blend * 0.5, 0.0001);

    // The sphere: D2 and D3, a few cents apart in each ear, brighter as it rises.
    float hs = spherePos(ts, 1.0).y / 0.22 * 0.5 + 0.5;
    float bs = 0.2 + 0.6 * hs * min(p.lift, 1.0);
    float2 low = float2(voice(73.42 * 1.001 * t, bs) + 0.6 * voice(146.83 * 0.999 * t, bs),
                        voice(73.42 * 0.999 * t, bs) + 0.6 * voice(146.83 * 1.001 * t + 0.25, bs)) * 0.1;

    // The ring: A3, swelling and shimmering as it comes close enough to the sphere to melt into it.
    float hr = ringPos(ts, 1.0).y / 0.22 * 0.5 + 0.5;
    float br = 0.2 + 0.6 * hr * min(p.lift, 1.0);
    float gap = abs(ringPos(ts, p.lift).y - spherePos(ts, p.lift).y) + 0.35 * abs(sin(ts * 0.7));
    float melt = exp(-gap * (3.0 / (0.2 + k * 4.0)));
    float shimmer = 1.0 + 0.5 * melt * sin(tau * 6.0 * t);
    float ring = voice(220.0 * t, br) * (0.35 + 0.65 * melt) * shimmer * 0.06;

    // The box: F#4, panned to its place on the screen and louder as it comes towards the camera.
    float3 ro = cameraPos(ts);
    float3 fw = normalize(-ro);
    float3 rt = normalize(cross(fw, float3(0.0, 1.0, 0.0)));
    float3 bp = boxPos(ts, p.lift);
    float3 rel = bp - ro;
    float side = clamp(dot(rel, rt) / max(dot(rel, fw), 0.5) * 0.9, -1.0, 1.0);
    float near = 1.0 / (1.0 + 0.25 * dot(rel, rel));
    float hb = 0.5 + 0.5 * sin(ts * 1.7);
    float box = voice(369.99 * t, 0.2 + 0.6 * hb * min(p.lift, 1.0)) * near * 0.16;
    float pan = 0.5 + 0.4 * side;
    float2 high = float2(1.0 - pan, pan) * box * 1.4;

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