import { AudioPlaybackContext } from '@/lib/useAudioPlayback';
import styled from '@emotion/styled';
import { useCallback, useContext, useEffect, useMemo, useRef } from 'react';

function initShaderProgram(gl: WebGLRenderingContext, vsSource: string, fsSource: string) {
  const vertexShader = loadShader(gl, gl.VERTEX_SHADER, vsSource)!;
  const fragmentShader = loadShader(gl, gl.FRAGMENT_SHADER, fsSource)!;

  const shaderProgram = gl.createProgram()!;
  gl.attachShader(shaderProgram, vertexShader);
  gl.attachShader(shaderProgram, fragmentShader);
  gl.linkProgram(shaderProgram);

  if (!gl.getProgramParameter(shaderProgram, gl.LINK_STATUS)) {
    console.log(
      `Unable to initialize the shader program: ${gl.getProgramInfoLog(
        shaderProgram,
      )}`,
    );
    return null;
  }

  return shaderProgram;
}

function loadShader(gl: WebGLRenderingContext, type: number, source: string) {
  const shader = gl.createShader(type)!;

  gl.shaderSource(shader, source);
  gl.compileShader(shader);

  if (!gl.getShaderParameter(shader, gl.COMPILE_STATUS)) {
    console.log(
      `An error occurred compiling the shaders: ${gl.getShaderInfoLog(shader)}`,
    );
    gl.deleteShader(shader);
    return null;
  }

  return shader;
}

const getLocations = (gl: WebGLRenderingContext, program: WebGLProgram, attribs: string[], uniforms: string[]) => {
  const res: {[key: string]: number} = {}
  for(const v of attribs) {
    res[v] = gl.getAttribLocation(program, v)!
  }
  for(const v of uniforms) {
    res[v] = gl.getUniformLocation(program, v)! as number;
  }
  return res
}

const setupWebGL = (gl: WebGLRenderingContext) => {
  const vertices = [-1, -1, -1, 1, 1, -1, -1, 1, 1, -1, 1, 1];
  const vertexBuffer = gl.createBuffer();
  gl.bindBuffer(gl.ARRAY_BUFFER, vertexBuffer);
  gl.bufferData(gl.ARRAY_BUFFER, new Float32Array(vertices), gl.STATIC_DRAW);

  const vsSource = `
    attribute vec2 coordinates;
    varying vec2 vPos;

    void main(void) {
      vPos = coordinates;
      gl_Position = vec4(vPos, 1.0, 1.0);
    }
  `
  const fsSource = `
    #define PI 3.1415926538

    precision highp float;
    varying vec2 vPos;
    uniform float a;
    uniform float n;
    uniform float m;
    uniform float b;
    uniform float t;
    uniform float s;

    float PHI = 1.61803398874989484820459;  // Φ = Golden Ratio

    float rand(in vec2 xy, in float seed){
          return fract(tan(distance(xy*PHI, xy)*seed)*xy.x);
    }

    void main(void) {

      float zr = a * sin(PI * n * (vPos.x)) * sin(PI * m * (vPos.y)) + b * sin(PI * m * (vPos.x)) * sin(PI * n * (vPos.y));

      float alphaR = pow((abs(zr)) / 2.0, 2.0) * 0.3;

      float rt = PI * 2.0 * (rand(vec2(vPos.x * t * 10000.0, vPos.y * t * 10000.0), 1.0));
      float rx = alphaR * cos(rt);
      float ry = alphaR * sin(rt);

      float z = a * sin(PI * n * (vPos.x + rx)) * sin(PI * m * (vPos.y + ry)) + b * sin(PI * m * (vPos.x + rx)) * sin(PI * n * (vPos.y + ry));

      float thr = pow((1.0 - (sqrt(pow(vPos.x, 2.0) + pow(vPos.y, 2.0)) * 0.99)) * 0.08, 1.25);

      float d = pow(abs(z) / thr, 1.5);

      float targetAlpha = (1.0 / d);
      float alpha = targetAlpha > (rt + 0.1) ? 1.0 : targetAlpha;
      float finalAlpha = s * min(1.0, d > 0.0 ? alpha : 0.0);

      gl_FragColor = vec4(
        0.361 + ((1.0 - 0.361) * finalAlpha),
        0.212 + ((1.0 - 0.212) * finalAlpha),
        0.682 + ((1.0 - 0.682) * finalAlpha),
        1.0
      );
    }
  `;

  const shaderProgram = initShaderProgram(gl, vsSource, fsSource);

  return {
    vertexBuffer,
    shaderProgram,
    locations: getLocations(gl, shaderProgram!, ["coordinates"], ["a", "n", "m", "b", "t", "s"]),
  };
}

const Wrapper = styled.div`
  height: 100%;
  display: flex;
  position: relative;
`;

const Canvas = styled.canvas`
  position: absolute;
  top: 0;
  left: 0;
  right: 0;
  bottom: 0;
`;

const FAST_HALF_LIFE = 0.250;

const energy: number[] = []

const TrackVisualizer = ({ bands }: { bands: number[][] }) => {

  const { playing, duration, getCurrentTime } = useContext(AudioPlaybackContext)

  const scaledBands = useMemo(() => {
    return bands.map((b) => {
      const squared = b.map((v) => v * v);
      const avg = squared.reduce((a, v) => a + v, 0) / squared.length;
      return squared.map((v) => v / avg);
    })
  }, [bands]);

  const canvasRef = useRef<HTMLCanvasElement | null>(null);
  const glRef = useRef<WebGLRenderingContext | null>(null);
  const glCtxRef = useRef<ReturnType<typeof setupWebGL> | null>(null);
  const playingRef = useRef<boolean>(false);

  useEffect(() => {
    const handleResize = () => {
      const canvas = canvasRef.current;
      if (!canvas) return;
      canvas.width = Math.round(canvas.parentElement!.clientWidth / 2) * 2;
      canvas.height = Math.round(canvas.parentElement!.clientHeight / 2) * 2;
    }
    window.addEventListener('resize', handleResize);
  }, []);

  const startPlaying = useCallback(() => {
    const canvas = canvasRef.current;
    const gl = glRef.current;
    const glCtx = glCtxRef.current;
    if (!canvas || !gl || !glCtx) return () => {};

    const { shaderProgram, vertexBuffer, locations } = glCtx;

    canvas.width = Math.round(canvas.parentElement!.clientWidth / 2) * 2;
    canvas.height = Math.round(canvas.parentElement!.clientHeight / 2) * 2;

    const bandLength = scaledBands[0].length;

    const frame = (frameDelta: number) => {

      const time = getCurrentTime();
      const delta = frameDelta;
      const decayPow = delta / FAST_HALF_LIFE;

      const progress = time / duration;
      const index = Math.floor(progress * bandLength);

      scaledBands.forEach((b, i) => {
        if (!energy[i]) energy[i] = 0;
        energy[i] -= delta / 2;
        energy[i] *= (Math.pow(0.25, decayPow));
        if (playingRef.current) {
          energy[i] += b[index] * 0.05;
        }
      });

      const minD = Math.min(canvas.width, canvas.height);
      const wOffset = Math.round((canvas.width - minD) / 2);
      const hOffset = Math.round((canvas.height - minD) / 2);

      gl.viewport(wOffset, hOffset, minD, minD);
      gl.clearColor(92/255, 54/255, 174/255, 1.0);
      gl.clear(gl.COLOR_BUFFER_BIT);

      const s = (Math.max(energy[0], energy[1], energy[2], energy[3]) - 0.01) * 20
      if (s <= 0) return;

      gl.useProgram(shaderProgram);
      gl.bindBuffer(gl.ARRAY_BUFFER, vertexBuffer);
      gl.vertexAttribPointer(locations.coordinates as number, 2, gl.FLOAT, false, 0, 0);
      gl.enableVertexAttribArray(locations.coordinates as number);

      gl.uniform1f(locations.t, Math.random());
      gl.uniform1f(locations.a, -2 * energy[2]);
      gl.uniform1f(locations.b, 4 * energy[3]);
      gl.uniform1f(locations.n, 1.05 + energy[0] + energy[2]);
      gl.uniform1f(locations.m, 0.65 + energy[1] + energy[3]);
      gl.uniform1f(locations.s, s);
      gl.drawArrays(gl.TRIANGLE_STRIP, 0, 6);
    }

    let stopped = false;

    let lastFrameTime = Date.now();
    const frameLoop = (restart?: boolean) => {
      if (!stopped) {
        const now = Date.now();
        frame((now - lastFrameTime) / 1000);
        lastFrameTime = now;
        requestAnimationFrame(() => frameLoop());
      }
    }

    frameLoop(true);

    return () => stopped = true;
  }, [duration, scaledBands, getCurrentTime]);

  const receiveCanvasRef = useCallback((canvas: HTMLCanvasElement | null) => {
    if (canvas) {
      canvasRef.current = canvas;
      glRef.current = canvas.getContext('webgl', { premultipliedAlpha: true })!;
      glCtxRef.current = setupWebGL(glRef.current);
    } else {
      canvasRef.current = null;
      glRef.current = null;
      glCtxRef.current = null;
    }
  }, [bands, playing]);

  useEffect(() => {
    return startPlaying();
  }, [startPlaying]);

  useEffect(() => {
    playingRef.current = playing;
  }, [playing]);

  return (
    <Wrapper>
      <Canvas ref={receiveCanvasRef} />
    </Wrapper>
  )
}

export default TrackVisualizer;
