
#include <metal_stdlib>
using namespace metal;

constant float PI = 3.14159265358979323846;

struct Uniforms_MaskingDiamondGradient {
    float scale;
    float center_offset;
};

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

float masking_diamond_gradient_mod(float a, float b) {
    return a - b * floor(a / b);
}

float2 masking_diamond_gradient_rotate_point(float2 point, float angle) {
    float cosTheta = cos(angle);
    float sinTheta = sin(angle);
    float2x2 rotationMatrix = float2x2(
        cosTheta, -sinTheta,
        sinTheta,  cosTheta
    );
    return rotationMatrix * point;
}

fragment float4 masking_diamond_gradient_frag(
VertexOut_MaskingDiamondGradient in [[ stage_in ]],
constant Uniforms_MaskingDiamondGradient &u [[ buffer(0) ]],
float4 color [[ color(0) ]]) {
    
    float2 uv = in.uv;
    float scale = u.scale;
    float center_offset = u.center_offset;
    
    uv -= 0.5;
    uv *= 2.0;
    uv *= scale;
    
//    uv.x -= (x_offset * scale);
//    uv.y += (y_offset * scale);
    
    // UV for diagonal from top left to bottom right and bottom left to top right
    float angle = (PI * 2.0) * 0.125;
    float2 uvTLBR = masking_diamond_gradient_rotate_point(uv, angle);
    float yValueTLBR = abs(uvTLBR.y);
    yValueTLBR -= center_offset;
    yValueTLBR = mix(yValueTLBR, 0.0, step(1.0, yValueTLBR));
    
    float2 uvBLTR = masking_diamond_gradient_rotate_point(uv, -angle);
    float yValueBLTR = abs(uvBLTR.y);
    yValueBLTR -= center_offset;
    yValueBLTR = mix(yValueBLTR, 0.0, step(1.0, yValueBLTR));
    
    float dist = max(yValueBLTR, yValueTLBR);
    
    return float4(dist, dist, dist, 1.0);
}
