#include <iostream>
#include <cmath>
#include <GL/glew.h>
#include <GLFW/glfw3.h>
#include <glm.hpp>
#include <gtc/type_ptr.hpp>
#include <vec3.hpp>
#include <vec4.hpp>
#include <mat4x4.hpp>
#include "MT3D.h"

#define ASSERT(x) if (!(x)) __debugbreak();
#define GLCall(x) GLClearError();x;ASSERT(GLLogCall(#x, __LINE__));

using namespace std;
using namespace glm;

static void GLClearError() {
    while (glGetError() != GL_NO_ERROR);
}

static bool GLLogCall(const char* function, int line) {
    while (GLenum error = glGetError()) {
        cout << "Error (" << error << ") : " << function << " - line " << line << endl;
        return false;
    }
    return true;
}

static GLuint compileShader(GLuint type, const string& source) {
    GLCall(GLuint id = glCreateShader(type));
    const char* src = source.c_str();

    GLCall(glShaderSource(id, 1, &src, nullptr));
    GLCall(glCompileShader(id));

    return id;
}

static int createShader(const string& vertexShader, const string& fragmentShader) {
    GLCall(GLuint program = glCreateProgram());
    GLCall(GLuint vs = compileShader(GL_VERTEX_SHADER, vertexShader));
    GLCall(GLuint fs = compileShader(GL_FRAGMENT_SHADER, fragmentShader));

    GLCall(glAttachShader(program, vs));
    GLCall(glAttachShader(program, fs));
    GLCall(glLinkProgram(program));
    GLCall(glValidateProgram(program));

    GLCall(glDeleteShader(vs));
    GLCall(glDeleteShader(fs));

    return program;
}

string vertexShader = "#version 300 es\n"
    "in vec4 a_vrhXYZ;\n"
    "in vec3 a_normala;\n"
    "uniform mat4 u_mTrans;\n"
    "uniform vec3 izvor;\n"
    "uniform vec3 kameraXYZ;\n"
    "out float svjetlina;\n"
    "\n"
    "void main() {\n"
    "   vec4 vrh = u_mTrans * a_vrhXYZ;\n"
    "   vec3 normala = mat3(u_mTrans) * a_normala;\n"
    "   vec3 premaIzvoru = normalize(izvor - vec3(vrh));\n"
    "   svjetlina = dot(premaIzvoru, normala);\n"
    "   float refleksija = 0.0;\n"
    "   if (svjetlina > 0.0) {\n"
    "       vec3 premaKameri = normalize(kameraXYZ - vec3(vrh));\n"
    "       vec3 reflektiranaZraka = reflect(-premaIzvoru, normala);\n"
    "       refleksija = max(dot(reflektiranaZraka, premaKameri), 0.0);\n"
    "       refleksija = pow(refleksija, 8.0);\n"
    "   }\n"
    "   svjetlina = svjetlina * 1.0 + refleksija * 0.0;\n"
    "   gl_Position = vrh;\n"
    "}\n";

string fragmentShader = "#version 300 es\n"
    "precision highp float;\n"
    "out vec4 bojaPiksela;\n"
    "in float svjetlina;\n"
    "uniform vec3 u_boja;\n"
    "\n"
    "void main() {\n"
    "   bojaPiksela = vec4(vec3(u_boja.x, u_boja.y, u_boja.z) * svjetlina, 1);\n"
    "}\n";

float* roller(float r, float h, int n) {
    const int len = 6 * (2 + 4 * (n + 1));
    float *vrhovi = new float[len];
    int j = 0;

    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = h / 2;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    float fi = 2 * PI / n;
    for (int i = 0; i <= n; i++) {
        float x = r * cos(fi);
        float y = r * sin(fi);

        vrhovi[j++] = x;
        vrhovi[j++] = y;
        vrhovi[j++] = h / 2;
        vrhovi[j++] = 0;
        vrhovi[j++] = 0;
        vrhovi[j++] = 1;

        fi += 2 * PI / n;
    }

    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = -h / 2;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = -1;

    fi = 2 * PI;

    for (int i = 0; i <= n; i++) {
        float x = r * cos(fi);
        float y = r * sin(fi);

        vrhovi[j++] = x;
        vrhovi[j++] = y;
        vrhovi[j++] = -h / 2;
        vrhovi[j++] = 0;
        vrhovi[j++] = 0;
        vrhovi[j++] = -1;

        fi -= 2 * PI / n;
    }

    fi = 0;

    for (int i = 0; i <= n; i++) {
        float C = cos(fi);
        float s = sin(fi);
        float x = r * C;
        float y = r * s;

        vrhovi[j++] = x;
        vrhovi[j++] = y;
        vrhovi[j++] = h / 2;
        vrhovi[j++] = C;
        vrhovi[j++] = s;
        vrhovi[j++] = 0;
        
        vrhovi[j++] = x;
        vrhovi[j++] = y;
        vrhovi[j++] = -h / 2;
        vrhovi[j++] = C;
        vrhovi[j++] = s;
        vrhovi[j++] = 0;

        fi += 2 * PI / n;
    }

    return vrhovi;
}

float* cone(float h, float r, float d) {
    float len = 6 * (d * 4 + 4);
    float* vrhovi = new float[len];
    int j = 0;

    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = h;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    for (float i = 0; i <= PI * 2; i += PI / d) {
        vrhovi[j++] = r * cos(i);
        vrhovi[j++] = r * sin(i);
        vrhovi[j++] = 0;
        vrhovi[j++] = cos(i);
        vrhovi[j++] = sin(i);
        vrhovi[j++] = 0;
    }

    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = -1;

    for (float i = 0; i <= PI * 2; i += PI / d) {
        vrhovi[j++] = r * cos(i);
        vrhovi[j++] = r * sin(i);
        vrhovi[j++] = 0;
        vrhovi[j++] = cos(i);
        vrhovi[j++] = sin(i);
        vrhovi[j++] = 0;
    }

    return vrhovi;
}

float* surface(float d) {
    float* vrhovi = new float[36];
    int j = 0;

    vrhovi[j++] = d;
    vrhovi[j++] = d;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    vrhovi[j++] = -d;
    vrhovi[j++] = -d;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    vrhovi[j++] = d;
    vrhovi[j++] = -d;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    vrhovi[j++] = -d;
    vrhovi[j++] = d;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    vrhovi[j++] = -d;
    vrhovi[j++] = -d;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    vrhovi[j++] = d;
    vrhovi[j++] = d;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 0;
    vrhovi[j++] = 1;

    return vrhovi;
}

float* ball(float r, int n) {
    float* vrhovi = new float[36 * 2 * n * n];
    int k = 0;

    float x, y, z;
    float xp = PI / (float) n;
    float yp = xp;

    for (float i = 0; i < PI; i += yp) {
        for (float j = 0; j < PI; j += xp) {
            x = cos(j + xp) * sin(i + yp);
            y = sin(j + xp) * sin(i + yp);
            z = cos(i + yp);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = -x;
            vrhovi[k++] = -y;
            vrhovi[k++] = -z;
            
            x = cos(j) * sin(i);
            y = sin(j) * sin(i);
            z = cos(i);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = -x;
            vrhovi[k++] = -y;
            vrhovi[k++] = -z;

            x = cos(j + xp) * sin(i);
            y = sin(j + xp) * sin(i);
            z = cos(i);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = -x;
            vrhovi[k++] = -y;
            vrhovi[k++] = -z;

            x = cos(j) * sin(i);
            y = sin(j) * sin(i);
            z = cos(i);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = -x;
            vrhovi[k++] = -y;
            vrhovi[k++] = -z;

            x = cos(j + xp) * sin(i + yp);
            y = sin(j + xp) * sin(i + yp);
            z = cos(i + yp);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = -x;
            vrhovi[k++] = -y;
            vrhovi[k++] = -z;

            x = cos(j) * sin(i + yp);
            y = sin(j) * sin(i + yp);
            z = cos(i + yp);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = -x;
            vrhovi[k++] = -y;
            vrhovi[k++] = -z;
        }
    }

    for (float i = 0; i < PI; i += yp) {
        for (float j = 0; j < PI; j += xp) {
            x = cos(j) * sin(i);
            y = sin(j) * sin(i);
            z = cos(i);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = x;
            vrhovi[k++] = y;
            vrhovi[k++] = z;

            x = cos(j + xp) * sin(i + yp);
            y = sin(j + xp) * sin(i + yp);
            z = cos(i + yp);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = x;
            vrhovi[k++] = y;
            vrhovi[k++] = z;

            x = cos(j + xp) * sin(i);
            y = sin(j + xp) * sin(i);
            z = cos(i);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = x;
            vrhovi[k++] = y;
            vrhovi[k++] = z;

            x = cos(j + xp) * sin(i + yp);
            y = sin(j + xp) * sin(i + yp);
            z = cos(i + yp);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = x;
            vrhovi[k++] = y;
            vrhovi[k++] = z;

            x = cos(j) * sin(i);
            y = sin(j) * sin(i);
            z = cos(i);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = x;
            vrhovi[k++] = y;
            vrhovi[k++] = z;

            x = cos(j) * sin(i + yp);
            y = sin(j) * sin(i + yp);
            z = cos(i + yp);

            vrhovi[k++] = x * r;
            vrhovi[k++] = y * r;
            vrhovi[k++] = z * r;
            vrhovi[k++] = x;
            vrhovi[k++] = y;
            vrhovi[k++] = z;
        }
    }

    return vrhovi;
}

GLuint rollerObj;
GLuint surfaceObj;
GLuint coneObj;
GLuint ballObj;

void fillContainers(GLuint shader) {
    GLCall(int a_vrhXYZ = glGetAttribLocation(shader, "a_vrhXYZ"));
    GLCall(int a_normala = glGetAttribLocation(shader, "a_normala"));
    
    GLuint rollerBuff;
    GLCall(glCreateVertexArrays(1, &rollerObj));
    GLCall(glCreateBuffers(1, &rollerBuff));
    GLCall(glBindVertexArray(rollerObj));
    GLCall(glBindBuffer(GL_ARRAY_BUFFER, rollerBuff));
    GLCall(glEnableVertexAttribArray(a_vrhXYZ));
    GLCall(glEnableVertexAttribArray(a_normala));
    GLCall(glVertexAttribPointer(a_vrhXYZ, 3, GL_FLOAT, GL_FALSE, 24, 0));
    GLCall(glVertexAttribPointer(a_normala, 3, GL_FLOAT, GL_FALSE, 24, (const void *)12));
    float r = 0.8, h = 0.8;
    int n = 32;
    float* rollerVerts = roller(r, h, n);
    int len = 6 * (2 + 4 * (n + 1));
    GLCall(glBufferData(GL_ARRAY_BUFFER, len * sizeof(float), rollerVerts, GL_STATIC_DRAW));

    GLuint surfaceBuff;
    GLCall(glCreateVertexArrays(1, &surfaceObj));
    GLCall(glCreateBuffers(1, &surfaceBuff));
    GLCall(glBindVertexArray(surfaceObj));
    GLCall(glBindBuffer(GL_ARRAY_BUFFER, surfaceBuff));
    GLCall(glEnableVertexAttribArray(a_vrhXYZ));
    GLCall(glEnableVertexAttribArray(a_normala));
    GLCall(glVertexAttribPointer(a_vrhXYZ, 3, GL_FLOAT, GL_FALSE, 24, 0));
    GLCall(glVertexAttribPointer(a_normala, 3, GL_FLOAT, GL_FALSE, 24, (const void *)12));
    n = 6;
    float* surfaceVerts = surface(n);
    len = n * n;
    GLCall(glBufferData(GL_ARRAY_BUFFER, len * sizeof(float), surfaceVerts, GL_STATIC_DRAW));

    GLuint coneBuff;
    GLCall(glCreateVertexArrays(1, &coneObj));
    GLCall(glCreateBuffers(1, &coneBuff));
    GLCall(glBindVertexArray(coneObj));
    GLCall(glBindBuffer(GL_ARRAY_BUFFER, coneBuff));
    GLCall(glEnableVertexAttribArray(a_vrhXYZ));
    GLCall(glEnableVertexAttribArray(a_normala));
    GLCall(glVertexAttribPointer(a_vrhXYZ, 3, GL_FLOAT, GL_FALSE, 24, 0));
    GLCall(glVertexAttribPointer(a_normala, 3, GL_FLOAT, GL_FALSE, 24, (const void*)12));
    h = 5; 
    r = 2;
    n = 32;
    float* coneVerts = cone(h, r, n);
    len = 6 * (n * 4 + 4);
    GLCall(glBufferData(GL_ARRAY_BUFFER, len * sizeof(float), coneVerts, GL_STATIC_DRAW));
    
    GLuint ballBuff;
    GLCall(glCreateVertexArrays(1, &ballObj));
    GLCall(glCreateBuffers(1, &ballBuff));
    GLCall(glBindVertexArray(ballObj));
    GLCall(glBindBuffer(GL_ARRAY_BUFFER, ballBuff));
    GLCall(glEnableVertexAttribArray(a_vrhXYZ));
    GLCall(glEnableVertexAttribArray(a_normala));
    GLCall(glVertexAttribPointer(a_vrhXYZ, 3, GL_FLOAT, GL_FALSE, 24, 0));
    GLCall(glVertexAttribPointer(a_normala, 3, GL_FLOAT, GL_FALSE, 24, (const void*)12));
    n = 32;
    r = 1;
    float* ballVertsU = ball(r, n);
    len = 36 * 2 * n * n;
    GLCall(glBufferData(GL_ARRAY_BUFFER, len * sizeof(float), ballVertsU, GL_STATIC_DRAW));    
}

int u_mTrans;
int izvor;
int kameraXYZ;
int u_boja;
MT3D *mt3d;

void drawRoller() {
    arr polje = mt3d->toList();
    mat4 umt = make_mat4x4(polje.list);
    GLCall(glUniformMatrix4fv(u_mTrans, 1, GL_FALSE, value_ptr(umt)));
    GLCall(glBindVertexArray(rollerObj));
    int n = 32;
    GLCall(glDrawArrays(GL_TRIANGLE_FAN, 0, n + 2));
    GLCall(glDrawArrays(GL_TRIANGLE_FAN, n + 2, n + 2));
    GLCall(glDrawArrays(GL_TRIANGLE_STRIP, 2 * (n + 2), 2 * n + 2));
}

void drawSurface() {
    arr polje = mt3d->toList();
    mat4 umt = make_mat4x4(polje.list);
    GLCall(glUniformMatrix4fv(u_mTrans, 1, GL_FALSE, value_ptr(umt)));
    GLCall(glBindVertexArray(surfaceObj));
    GLCall(glDrawArrays(GL_TRIANGLES, 0, 36));
}

void drawCone() {
    arr polje = mt3d->toList();
    mat4 umt = make_mat4x4(polje.list);
    GLCall(glUniformMatrix4fv(u_mTrans, 1, GL_FALSE, value_ptr(umt)));
    GLCall(glBindVertexArray(coneObj));
    int n = 32;
    GLCall(glDrawArrays(GL_TRIANGLE_FAN, 0, 2 * n + 2));
    GLCall(glDrawArrays(GL_TRIANGLE_FAN, 2 * n + 2, 2 * n + 2));
}

void drawBall() {
    arr polje = mt3d->toList();
    mat4 umt = make_mat4x4(polje.list);
    GLCall(glUniformMatrix4fv(u_mTrans, 1, GL_FALSE, value_ptr(umt)));
    int n = 32; 
    GLCall(glBindVertexArray(ballObj));
    GLCall(glDrawArrays(GL_TRIANGLES, 0, 36 * 2 * n * n));
}

void reset() {
    mt3d->identity();
    mt3d->persp(-2, 2, -2, 2, 2.6, 20);
    mt3d->setCamera(6, 6, 8, 0, 0, 0, 0, 0, 0.1);
}

int main() {
    GLFWwindow* window;
    mt3d = new MT3D;

    if (!glfwInit()) {
        return -1;
    }

    window = glfwCreateWindow(960, 960, "OpenGl", NULL, NULL);
    if (!window) {
        glfwTerminate();
        return -1;
    }

    glfwMakeContextCurrent(window);
    
    glewInit();
    cout << glGetString(GL_VERSION) << endl;

    GLCall(glEnable(GL_CULL_FACE));
    GLCall(glEnable(GL_DEPTH_TEST));

    GLuint shader = createShader(vertexShader, fragmentShader);
    GLCall(glUseProgram(shader));

    GLCall(u_mTrans = glGetUniformLocation(shader, "u_mTrans"));
    GLCall(izvor = glGetUniformLocation(shader, "izvor"));
    GLCall(kameraXYZ = glGetUniformLocation(shader, "kameraXYZ"));
    GLCall(u_boja = glGetUniformLocation(shader, "u_boja"));

    fillContainers(shader);

    float rot = 0;

    while (!glfwWindowShouldClose(window)) {

        GLCall(glClearColor(0, 0, 0, 1));
        GLCall(glClear(GL_COLOR_BUFFER_BIT | GL_DEPTH_BUFFER_BIT));
        GLCall(glViewport(0, 0, 960, 960));
        GLCall(glUniform3f(izvor, -5, 3, -6));
        GLCall(glUniform3f(kameraXYZ, 6, 6, 8));

        reset();
        GLCall(glUniform3f(u_boja, 0, 0.8, 0.8)); //cyan
        drawSurface();
        GLCall(glUniform3f(u_boja, 0.8, 0.5, 0)); //brown
        drawCone();
        mt3d->translate(0, 0, 3.5);
        drawRoller();

        float sx = 0.4, sy = 0.4, sz = 6;
        float px = 0, py = 8.5, pz = 0.4;
        GLCall(glUniform3f(u_boja, 0.6, 0.6, 0)); //yellow
        reset();
        mt3d->rotateZ(rot);
        mt3d->rotateX(90);
        mt3d->scale(sx, sy, sz);
        mt3d->translate(px, py, pz);
        drawRoller();
        reset();
        mt3d->rotateZ(rot - 120);
        mt3d->rotateX(90);
        mt3d->scale(sx, sy, sz);
        mt3d->translate(px, py, pz);
        drawRoller();
        reset();
        mt3d->rotateZ(rot - 240);
        mt3d->rotateX(90);
        mt3d->scale(sx, sy, sz);
        mt3d->translate(px, py, pz);
        drawRoller();

        px = 0;
        py = 4.6;
        pz = 3.6;
        GLCall(glUniform3f(u_boja, 0.8, 0.8, 0)); //yellow
        reset();
        mt3d->rotateZ(rot - 60);
        mt3d->translate(px, py, pz);
        mt3d->rotateZ(90);
        drawBall();
        reset();
        mt3d->rotateZ(rot - 180);
        mt3d->translate(px, py, pz);
        mt3d->rotateZ(90);
        drawBall();
        reset();
        mt3d->rotateZ(rot - 300);
        mt3d->translate(px, py, pz);
        mt3d->rotateZ(90);
        drawBall();

        rot += 0.02;
        if (rot >= 360) {
            rot = 0;
        }

        glfwSwapBuffers(window);
        glfwPollEvents();
    }

    glfwTerminate();
    return 0;
}