#include <stdio.h>
#include <stdlib.h>
#include "mpi.h"

#define MEGA_BYTE (1024 * 1024)

int main(int argc, char *argv[])
{
  int rank, size, i, buf_size_in_mb, destrank, rep;
  char *buffer = NULL, *dummy;
  double start, end, duration, dura_sum, dura_max;
  MPI_Win win;

  if (argc != 3) {
    printf("Usage: mpiexec -n np %s BUF_SIZE REP\n", argv[0]);
    exit(-1);
  }

  buf_size_in_mb = atoi(argv[1]);
  rep = atoi(argv[2]);

  MPI_Init(&argc, &argv);
  MPI_Comm_size(MPI_COMM_WORLD, &size);
  MPI_Comm_rank(MPI_COMM_WORLD, &rank);

  if (size % 2 != 0) {
    printf("Please run with even number of processes.\n"); fflush(stdout);
    MPI_Finalize();
    return 0;
  }

  buffer = (char *)malloc(buf_size_in_mb * MEGA_BYTE);
  MPI_Barrier(MPI_COMM_WORLD);
  start = MPI_Wtime();
  if (rank % 2 == 0) {
	MPI_Win_allocate(0, 1, MPI_INFO_NULL, MPI_COMM_WORLD, &dummy, &win);
	destrank = rank + 1;
	MPI_Win_lock(MPI_LOCK_SHARED, destrank, MPI_MODE_NOCHECK, win);
	for (i = 0; i < rep; i++) {
	  MPI_Put(buffer, buf_size_in_mb * MEGA_BYTE, MPI_BYTE, destrank, 0,
			  buf_size_in_mb * MEGA_BYTE, MPI_BYTE, win);
	  MPI_Win_flush(destrank, win);
	}
	MPI_Win_unlock(destrank, win);
  } else {
    MPI_Win_allocate(buf_size_in_mb * MEGA_BYTE, 1,
					 MPI_INFO_NULL, MPI_COMM_WORLD, &buffer, &win);
  }
  MPI_Barrier(MPI_COMM_WORLD);
  end = MPI_Wtime();
  duration = end - start;
  MPI_Reduce(&duration, &dura_max, 1, MPI_DOUBLE, MPI_MAX, 0, MPI_COMM_WORLD);
  if (rank == 0) {
    printf("Concurrent %d pairs of %dMB MPI_Put (%d repetitions) takes: "
           "%f seconds, aggregate bandwidth: %f GB/s\n", size / 2,
           buf_size_in_mb, rep, dura_max,
           buf_size_in_mb * size * rep / (dura_max * 2048));
  }

  if (rank % 2 == 0)
    if (buffer) free(buffer);
  MPI_Win_free(&win);
  MPI_Finalize();

  return 0;
}
