#include <mpi.h>
#include <omp.h>
#include <stdlib.h>
#include <stdio.h>
#include <string.h>

struct graph {
  int num_verts;
  int num_edges;
  int* out_array;
  int* out_degree_list;
} ;
#define out_degree(g, n) (g->out_degree_list[n+1] - g->out_degree_list[n])
#define out_vertices(g, n) &g->out_array[g->out_degree_list[n]]

int rank, size;

void read_edge(char* filename,
  int& num_verts, int& num_edges,
  int*& srcs, int*& dsts)
{
  double timer = omp_get_wtime();
  printf("Reading %s\n", filename);
  
  FILE* infile = fopen(filename, "r");
  char line[256];

  num_verts = 0;

  int count = 0;
  int cur_size = 1024*1024;

  srcs = (int*)malloc(cur_size*sizeof(int));
  dsts = (int*)malloc(cur_size*sizeof(int));

  while(fgets(line, 256, infile) != NULL) {
    if (line[0] == '%') continue;
    sscanf(line, "%d %d", &srcs[count], &dsts[count]);

    if (srcs[count] > num_verts)
      num_verts = srcs[count];
    if (dsts[count] > num_verts)
      num_verts = dsts[count];

    count += 1;
    if (count+1 > cur_size) {
      cur_size *= 2;
      srcs = (int*)realloc(srcs, cur_size*sizeof(int));
      dsts = (int*)realloc(dsts, cur_size*sizeof(int));
    }
  }  
  num_verts += 1;
  num_edges = count*2;
  
  fclose(infile);
  
  timer = omp_get_wtime() - timer;
  printf("Time: %lf (s)\n", timer);

  return;
}

void create_csr(int num_verts, int num_edges, 
  int* srcs, int* dsts,
  int*& out_array, int*& out_degree_list)
{
  double timer = omp_get_wtime();
  printf("Creating CSR\n");
  
  out_array = (int*)malloc(num_edges*sizeof(int));
  out_degree_list = (int*)malloc((num_verts+1)*sizeof(int));

  for (int i = 0; i < num_edges; ++i)
    out_array[i] = 0;
  for (int i = 0; i < num_verts+1; ++i)
    out_degree_list[i] = 0;

  int* temp_counts = (int*)malloc(num_verts*sizeof(int));
  for (int i = 0; i < num_verts; ++i)
    temp_counts[i] = 0;
  for (int i = 0; i < num_edges / 2; ++i) {
    ++temp_counts[srcs[i]];
    ++temp_counts[dsts[i]];
  }
  for (int i = 0; i < num_verts; ++i)
    out_degree_list[i+1] = out_degree_list[i] + temp_counts[i];
  memcpy(temp_counts, out_degree_list, num_verts*sizeof(int));
  for (int i = 0; i < num_edges / 2; ++i) {
    out_array[temp_counts[srcs[i]]++] = dsts[i];
    out_array[temp_counts[dsts[i]]++] = srcs[i];
  }
  
  free(temp_counts);
  
  timer = omp_get_wtime() - timer;
  printf("Time: %lf (s)\n", timer);
  
  return;
}


void do_bfs_serial(graph* g, int root)
{
  int* levels = (int*)malloc(g->num_verts*sizeof(int));
  int level = 0;

  for (int i = 0; i < g->num_verts; ++i)
    levels[i] = -1;

  levels[root] = level;

  int* queue = (int*)malloc(g->num_verts*sizeof(int));
  int* queue_next = (int*)malloc(g->num_verts*sizeof(int));
  int queue_size = 0;
  int next_size = 0;
  queue[queue_size++] = root;

  double timer = omp_get_wtime();
  printf("BFS Serial\n");
  while (queue_size > 0) {
    ++level;
    for (int i = 0; i < queue_size; ++i) {
      int vert = queue[i];
      int out_degree = out_degree(g, vert);
      int* out_edges = out_vertices(g, vert);
      for (int j= 0; j < out_degree; ++j) {
        int out = out_edges[j];
        if (levels[out] < 0) {
          levels[out] = level;
          queue_next[next_size++] = out;
        }
      }
    }

    int* temp = queue;
    queue = queue_next;
    queue_next = temp;
    queue_size = next_size;
    next_size = 0;
  }
  timer = omp_get_wtime() - timer;
  printf("Time: %lf (s)\n", timer);
  
  free(levels);
  free(queue);
  free(queue_next);
}



void do_bfs_parallel(graph* g, int root)
{
  int* levels = (int*)malloc(g->num_verts*sizeof(int));
  int level = 0;

  for (int i = 0; i < g->num_verts; ++i)
    levels[i] = -1;

  levels[root] = level;

  int* queue = (int*)malloc(g->num_verts*sizeof(int));
  int* queue_next = (int*)malloc(g->num_verts*sizeof(int));
  int queue_size = 0;
  int next_size = 0;
  queue[queue_size++] = root;

  double timer = omp_get_wtime();
  printf("BFS Parallel\n");
#pragma omp parallel
{
  int thread_queue[64];
  int thread_queue_size = 0;  

  while (queue_size > 0) {
  
#pragma omp single
{
    ++level;
}

#pragma omp for schedule(static)
    for (int i = 0; i < queue_size; ++i) {
      int vert = queue[i];
      int out_degree = out_degree(g, vert);
      int* out_edges = out_vertices(g, vert);
      for (int j= 0; j < out_degree; ++j) {
        int out = out_edges[j];
        if (levels[out] < 0) {
          levels[out] = level;
          thread_queue[thread_queue_size++] = out;
          if (thread_queue_size == 64) {
            int index;
      #pragma omp atomic capture
            index = next_size += 64;
            index -= 64;
            for (int k = 0; k < 64; ++k)
              queue_next[index+k] = thread_queue[k];
            thread_queue_size = 0;
          }
        }
      }
    }

    int index;
#pragma omp atomic capture
    index = next_size += thread_queue_size;
    index -= thread_queue_size;
    for (int k = 0; k < thread_queue_size; ++k)
      queue_next[index+k] = thread_queue[k];
    thread_queue_size = 0;
#pragma omp barrier

#pragma omp single
{
    int* temp = queue;
    queue = queue_next;
    queue_next = temp;
    queue_size = next_size;
    next_size = 0;
} // end single
  } // end while
} // end parallel
  timer = omp_get_wtime() - timer;
  printf("Time: %lf (s)\n", timer);

  free(levels);
  free(queue);
  free(queue_next);
}

int binary_search(int* offsets, int val, int bound_low, int bound_high)
{
  bool found = false;
  int index = 0;
  while (!found)
  {
    index = (bound_high + bound_low) / 2;
    if (offsets[index] <= val && offsets[index+1] > val)
    {
      found = true;
    }
    else if (offsets[index] <= val)
      bound_low = index;
    else if (offsets[index] > val)
      bound_high = index;
  }

  return index;
}

void do_bfs_loop_collapse(graph* g, int root)
{
  int* levels = (int*)malloc(g->num_verts*sizeof(int));
  int level = 0;

  int* queue = (int*)malloc(g->num_verts*sizeof(int));
  int* queue_next = (int*)malloc(g->num_verts*sizeof(int));
  int* offsets = (int*)malloc(g->num_verts*sizeof(int));
  int* offsets_next = (int*)malloc(g->num_verts*sizeof(int));
  int queue_size_verts = 0;
  int next_size_verts = 0;
  int queue_size_edges = out_degree(g, root);
  queue[queue_size_verts++] = root;
  offsets[0] = 0;
  offsets[1] = out_degree(g, root);


  double timer = omp_get_wtime();
  printf("BFS Loop collapse\n");
#pragma omp parallel
{
#pragma omp for
  for (int i = 0; i < g->num_verts; ++i)
    levels[i] = -1;

#pragma omp single
{
  levels[root] = level;
}

  int thread_queue[64];
  int thread_queue_size = 0;  

  while (queue_size_edges > 0) {
  
#pragma omp single
{
    ++level;
}

#pragma omp for schedule(static) nowait
    for (int i = 0; i < queue_size_edges; ++i) {
      int vert_index = binary_search(offsets, i, 0, queue_size_verts);
      int vert = queue[vert_index];

      int* outs = out_vertices(g, vert);
      int j = i - offsets[vert_index];
      int out = outs[j];

      if (levels[out] < 0) {
        levels[out] = level;
        thread_queue[thread_queue_size++] = out;
        if (thread_queue_size == 64) {
          int index;
    #pragma omp atomic capture
          index = next_size_verts += 64;
          index -= 64;
          for (int k = 0; k < 64; ++k) {
            queue_next[index+k] = thread_queue[k];
            offsets_next[index+k] = out_degree(g, thread_queue[k]);
          }
          thread_queue_size = 0;
        }
      }
    }

    int index;
#pragma omp atomic capture
    index = next_size_verts += thread_queue_size;
    index -= thread_queue_size;
    for (int k = 0; k < thread_queue_size; ++k) {
      queue_next[index+k] = thread_queue[k];
      offsets_next[index+k] = out_degree(g, thread_queue[k]);
    }
    thread_queue_size = 0;
#pragma omp barrier

#pragma omp single
{
    // could be parallelized - CUDA scan
    offsets[0] = 0;
    for (int i = 0; i < next_size_verts; ++i)
      offsets[i+1] = offsets[i] + offsets_next[i];
    queue_size_edges = offsets[next_size_verts];


    int* temp = queue;
    queue = queue_next;
    queue_next = temp;
    queue_size_verts = next_size_verts;
    next_size_verts = 0;
} 
  } // end while
} // end parallel
  timer = omp_get_wtime() - timer;
  printf("Time: %lf (s)\n", timer);

  free(levels);
  free(queue);
  free(queue_next);
}

void do_bfs_distributed(graph* g, int root)
{
  int* levels = (int*)malloc(g->num_verts*sizeof(int));
  int level = 0;

  for (int i = 0; i < g->num_verts; ++i)
    levels[i] = -1;

  int start_vert = rank * (g->num_verts / size + 1);
  int end_vert = (rank + 1) * (g->num_verts / size + 1);
  if (end_vert > g->num_verts)
    end_vert = g->num_verts;

  printf("Task: %d - start: %d - end: %d\n", rank, start_vert, end_vert);

  levels[root] = level;

  int* queue = (int*)malloc(g->num_verts*sizeof(int));
  int* queue_next = (int*)malloc(g->num_verts*sizeof(int));
  int queue_size = 0;
  int next_size = 0;
  queue[queue_size++] = root;

  // Set distances from root
  double timer = omp_get_wtime();
  printf("BFS Distributed\n");
  while (queue_size > 0) {
    ++level;
    for (int i = 0; i < queue_size; ++i) {
      int vert = queue[i];
      if (vert < start_vert || vert >= end_vert)
        continue;

      int out_degree = out_degree(g, vert);
      int* out_edges = out_vertices(g, vert);
      for (int j= 0; j < out_degree; ++j) {
        int out = out_edges[j];
        if (levels[out] < 0) {
          levels[out] = level;
          queue_next[next_size++] = out;
        }
      }
    }

    int* temp = queue;
    queue = queue_next;
    queue_next = temp;
    queue_size = next_size;
    next_size = 0;

    if (rank == 0) {
      MPI_Status status;
      MPI_Send(queue, queue_size, MPI_INT, 1, 0, MPI_COMM_WORLD);
      MPI_Recv(queue+queue_size, g->num_verts, MPI_INT, 1, 0, MPI_COMM_WORLD, &status);
      int recv_count = -1;
      MPI_Get_count(&status, MPI_INT, &recv_count); 
      queue_size += recv_count;
    }
    else
    {
      MPI_Status status;
      MPI_Recv(queue+queue_size, g->num_verts, MPI_INT, 0, 0, MPI_COMM_WORLD, &status);
      int recv_count = -1;
      MPI_Get_count(&status, MPI_INT, &recv_count); 
      MPI_Send(queue, queue_size, MPI_INT, 0, 0, MPI_COMM_WORLD);
      queue_size += recv_count;
    }
    MPI_Allreduce(MPI_IN_PLACE, levels, g->num_verts, MPI_INT, MPI_MAX, MPI_COMM_WORLD);
  }
  timer = omp_get_wtime() - timer;
  printf("Time: %lf (s)\n", timer);

  free(levels);
  free(queue);
  free(queue_next);
}


int main(int argc, char** argv)
{
  MPI_Init(&argc, &argv);
  MPI_Comm_rank(MPI_COMM_WORLD, &rank);
  MPI_Comm_size(MPI_COMM_WORLD, &size);
  
  int* srcs = NULL;
  int* dsts = NULL;
  int num_verts = 0;
  int num_edges = 0;
  int* out_array = NULL;
  int* out_degree_list = NULL;

  read_edge(argv[1], num_verts, num_edges, srcs, dsts);
  create_csr(num_verts, num_edges, srcs, dsts, 
    out_array, out_degree_list);
  graph g = {num_verts, num_edges, out_array, out_degree_list};
  free(srcs);
  free(dsts);

  int root = 0;

  do_bfs_serial(&g, root);
  do_bfs_parallel(&g, root);
  do_bfs_loop_collapse(&g, root);
  do_bfs_distributed(&g, root);

  free(out_array);
  free(out_degree_list);
  
  MPI_Barrier(MPI_COMM_WORLD);
  MPI_Finalize();

  return 0;
}
