// #include "mpi_test_incl.h"
#include <mpi.h>
#include <iostream>
#include <vector>
#include <string.h>
#include <stdlib.h>
#include <sys/time.h>

using namespace std;

const int BUF_SZ = 16 * 1024; // use 2Kbyte buffer for send & receive
const int TAG1 = 50;         // use 2 tags TAG1 & TAG2 and swap between them
const int freqSecReport = 20; // Report frequency.

using namespace std;

int         nIters        = 1000;
int         Scale         = 100000;
char        OutFilePrefix[256];

void PrintUsage(const char* ProgName);
void ReadParameters(int argc, char* argv[]);
void handleMPIerror(FILE* fpLog, const string str, int retErr, int rank, int iter, const MPI_Status* Status = NULL);

class Statistics
{
public:
    Statistics(int printFreqSec, int maxIter);
    void print(FILE* fp);
protected:
    int mPrintFreqSec;
    int mMaxIter;
    int mLastReportIter;
    int mCurIter;

    int mIter2Report;

    struct timeval mStartTime;
    struct timeval mLastReportTime;
};

int mainInit(int argc, char* argv[], int* clusterSize, int* procRank, FILE** fppLog);
void validateInput(FILE* fpLog, vector<char*>& RcvBufs, vector<int>& SlavesRcvIters, int SlaveIdx, Statistics& stat, int* SlavesFinished);

void Rcv_WaitAny(int Slaves, int iters, FILE* fpLog)
{
    Statistics  myStatistics(freqSecReport, iters);
    MPI_Status  status;
    int         slaveRank, slaveIdx;
    int         retErr  = MPI_SUCCESS;
    int         SlavesFinished = 0;  // number of slaves already finished

    vector<char*>       RcvBufs(Slaves);        // buffers for data from each slave
    vector<int>         SlavesRcvIters(Slaves); // count number of recv iterations per slave (init with 0)
    vector<MPI_Request> RcvRequests(Slaves);


    for(slaveIdx=0; slaveIdx<Slaves; ++slaveIdx)
    {
        RcvBufs[slaveIdx] = new char [BUF_SZ + 1];
        slaveRank = slaveIdx + 1;
        MPI_Irecv(RcvBufs[slaveIdx], BUF_SZ, MPI::CHAR, slaveRank, TAG1, MPI_COMM_WORLD, &(RcvRequests[slaveIdx]));
    }

    while(SlavesFinished < Slaves)
    {
        retErr = MPI_Waitany(RcvRequests.size(), &*RcvRequests.begin(), &slaveIdx, &status);
        slaveRank = slaveIdx + 1;

        handleMPIerror(fpLog, "Receiver recognized error", retErr, slaveRank, SlavesRcvIters[slaveIdx], &status);
        validateInput(fpLog, RcvBufs, SlavesRcvIters, slaveIdx, myStatistics, &SlavesFinished);


        MPI_Recv(RcvBufs[slaveIdx], BUF_SZ, MPI::CHAR, slaveRank, TAG1, MPI_COMM_WORLD, &status);
        validateInput(fpLog, RcvBufs, SlavesRcvIters, slaveIdx, myStatistics, &SlavesFinished);


        MPI_Irecv(RcvBufs[slaveIdx], BUF_SZ, MPI::CHAR, slaveRank, TAG1, MPI_COMM_WORLD, &(RcvRequests[slaveIdx]));
    }

    fprintf(fpLog, "\n\n\nEnd of master subroutine \n");
}


int SndSyncSlave(int rank, FILE* fpLog)
{
    Statistics myStatistics(freqSecReport, nIters);

    const int masterRank  = 0;

    char* chBuf = new char [BUF_SZ + 100];

    for(int i=1; i<nIters+1; ++i)
    {
        *(int*)chBuf = rank * Scale + i; // init with different number

        // send
        int retErr = MPI_Send(chBuf, BUF_SZ, MPI_CHAR, masterRank, TAG1, MPI_COMM_WORLD);
        handleMPIerror(fpLog, "Sender fail:", retErr, rank, i);
        myStatistics.print(fpLog);
    }

    fprintf(fpLog, "Last Sent number = %d\n\n\n\nEnd of slave subroutine\n", *(int*)chBuf);
    fflush(fpLog);
    return 0;
}




void validateInput(FILE* fpLog, vector<char*>& RcvBufs, vector<int>& SlavesRcvIters, int slaveIdx, Statistics& myStatistics, int* SlavesFinished)
{
    const int rank = slaveIdx + 1;
    const int expectedNum = rank * Scale + SlavesRcvIters[slaveIdx] + 1;
    int* pTestNum = (int*)RcvBufs[slaveIdx];

    if (expectedNum != *pTestNum) {
        fprintf(fpLog, "Validation error: From rank %d recevied %d expected %d\n", rank, *pTestNum, expectedNum);
    }
    ++SlavesRcvIters[slaveIdx]; // received one message
    myStatistics.print(fpLog);

    // test print
    if((SlavesRcvIters[slaveIdx] % 100000) == 0) {
        fprintf(fpLog, "\n\nFrom rank %d, Iters = %d\n ", rank, SlavesRcvIters[slaveIdx]);
    }

    if(SlavesRcvIters[slaveIdx] == nIters)
    {
        *SlavesFinished += 1;
        fprintf(fpLog, "\n\nSlave finished: rank %d, last number = %d Iters = %d\n ", rank, *pTestNum, SlavesRcvIters[slaveIdx]);
    }
    fflush(fpLog);
}



void PrintUsage(const char* ProgName)
{
    printf("Usage example: %s 10 100000 2 5 1 out\n", ProgName);
    printf("Where: \n");
    printf("%s - program exe file \n", ProgName);
    printf("10     - number of interactions between master & each one of slaves (= loop iterations in master & slaves)\n");
    printf("100000 - scale factor for integer sent by slave, slave sends 100000*rank + iteration number\n");
    printf("out    - prefix for output files. out_rank.txt files will be created.\n\n");
}

void ReadParameters(int argc, char* argv[])
{
    if(argc < 2)
    {
        PrintUsage(argv[0]);
        exit(-1);
    }

    int i = 1;
    sscanf(argv[i++], "%d", &nIters);
    sscanf(argv[i++], "%d", &Scale);
    strncpy(OutFilePrefix, argv[i++], 256);
}

void handleMPIerror(FILE* fpLog, const string str, int retErr, int rank, int iter, const MPI_Status* Status)
{
    if (retErr == MPI_SUCCESS)
        return;

    char  errString[MPI_MAX_ERROR_STRING];
    int   errStrLen;
    MPI_Error_string(retErr, errString, &errStrLen);

    fprintf(fpLog, "\n\n==========================================\n");
    fprintf(fpLog, "%s (rank = %d): \n", str.c_str(), rank);
    fprintf(fpLog, "------------------------------------------\n");
    fprintf(fpLog, "ErrorCode:    0x%X \n", retErr);
    fprintf(fpLog, "ErrorString:  %s \n"  , errString);
    if(Status != NULL)
    {
        // MPI_Error_string(Status->MPI_ERROR, errString, &errStrLen);
        fprintf(fpLog, "Status.MPI_ERROR(code) = %d \n"  , Status->MPI_ERROR);
        // fprintf(fpLog, "Status.MPI_ERROR(str)  = %s \n"  , errString);
        fprintf(fpLog, "Status.MPI_SOURCE      = %d \n"  , Status->MPI_SOURCE);
        fprintf(fpLog, "Status.MPI_TAG         = %d \n"  , Status->MPI_TAG);
        fprintf(fpLog, "Status.count           = %d \n"  , Status->count);
        fprintf(fpLog, "Status.cancelled       = %d \n"  , Status->cancelled);
    }
    fprintf(fpLog, "==========================================\n\n");
    fflush(fpLog);
}

Statistics::Statistics(int printFreqSec, int maxIter) :
mPrintFreqSec(printFreqSec), mMaxIter(maxIter), mLastReportIter(0), mCurIter(0), mIter2Report(0)
{
    gettimeofday(&mStartTime, NULL);
    mLastReportTime = mStartTime;
}

void Statistics::print(FILE* fp)
{
    ++mCurIter;

    if (mIter2Report > 0) { // avoid unusefull gettimeofday call, skip iterations estimated from previous report
        --mIter2Report;
        return;
    }

    struct timeval curTime;
    gettimeofday(&curTime, NULL);

    double curSecs = (curTime.tv_sec - mLastReportTime.tv_sec) + (curTime.tv_usec - mLastReportTime.tv_usec) / 1000000.0;

    if ( (curSecs) < mPrintFreqSec )
        return;

    mIter2Report = (mCurIter - mLastReportIter) / (curSecs / mPrintFreqSec);

    double msgSizeMb = BUF_SZ  * 8 / 1024.0 / 1024.0;
    double curRate = (mCurIter - mLastReportIter) / curSecs * msgSizeMb;
    double totSecs = (curTime.tv_sec - mStartTime.tv_sec) + (curTime.tv_usec - mStartTime.tv_usec) / 1000000.0;
    double totRate = mCurIter / totSecs * msgSizeMb;



    fprintf(fp, "Iter = %7d / %7d, Current rate: %6.2f Mbit/sec (for %4.1fsecs), Average rate: %6.2f Mbit/sec (for %4.1fsecs)\n"
              , mCurIter, mMaxIter, curRate, curSecs, totRate, totSecs);
    mLastReportTime = curTime;
    mLastReportIter = mCurIter;
}

void openOutFile(int rank, FILE** fppLog)
{
    char filename[256];
    sprintf(filename, "./%s_r%d.log", OutFilePrefix, rank);

    *fppLog = fopen(filename, "wt");
    fprintf(*fppLog, "rank = %d\n", rank);

    char hostname[256];
    gethostname(hostname, 256);
    fprintf(*fppLog, "hostname = %s\n", hostname);
}

int mainInit(int argc, char* argv[], int* clusterSize, int* procRank, FILE** fppLog)
{
    ReadParameters(argc, argv);

    MPI::Init_thread(MPI_THREAD_MULTIPLE);
    MPI_Errhandler_set(MPI_COMM_WORLD, MPI_ERRORS_RETURN); // return an exception instead of Fatal Error, which finishes execution of all processes.

    int  size, rank;

    MPI_Comm_size(MPI_COMM_WORLD, &size);
    MPI_Comm_rank(MPI_COMM_WORLD, &rank);

    openOutFile(rank, fppLog);

    *clusterSize = size;
    *procRank    = rank;

    return 0;
}


int main(int argc, char *argv[])
{
    int  size = 1, rank = 0;
    bool isAllSlavesLive = true;
    FILE* fpLog;
    unsigned long affinityMask    = 0x3;
    unsigned int  affinityMaskLen = sizeof(affinityMask);

    int retErr = mainInit(argc, argv, &size, &rank, &fpLog);
    if (retErr < 0)
        return 0;

    try
    {
        if(rank == 0)
        {
            Rcv_WaitAny(size-1, nIters, fpLog);
        }
        else
        {
            SndSyncSlave(rank, fpLog);
        }
    }

    catch(...)
    {
        fprintf(fpLog, "\n\n\nunexpected C++ exception\n\n");
    }

    bool isAbort = false;
    retErr = MPI_Finalize();
    if (retErr) {
        fprintf(fpLog, "Error in finalization");
    }
    fflush(fpLog);

    return 0;
}
