// #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 = 2 * 1024; // use 2Kbyte buffer for send & receive

const int TAG1 = 50;         // use 2 tags TAG1 & TAG2 and swap between them
// const int TAG2 = 50;
const int TAG2 = 100;

const int freqSecReport = 20; // Report frequency.

using namespace std;

int         nIters        = 1000;
int         Scale         = 100000;
int         failNodeRank  = 2;
int         failIteration = 5;
int         SleepSecs     = 2;
const int   EndOfSession  = 0;
char        OutFilePrefix[256];
bool        bExeSlaveBcast = false;
int         BcastNum       = -1;

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("2      - rank of fail node \n");
    printf("5      - Fail iteration. On this iteration, slave 2 will fail & the rest slaves will sleep 1 sec (see next parameter)\n");
    printf("1      - all slaves except failed one, will sleep 1 second on iteration of failure, to ensure that they finish after failure.\n");
    printf("         This mechanism achieves slave failure in the midle of test, but not after all slaves finished\n");
    printf("out    - prefix for output files. out_rank.txt files will be created.\n\n");
}

void ReadParameters(int argc, char* argv[])
{
    int i = 1;
    sscanf(argv[i++], "%d", &nIters);
    sscanf(argv[i++], "%d", &Scale);
    sscanf(argv[i++], "%d", &failNodeRank);
    sscanf(argv[i++], "%d", &failIteration);
    sscanf(argv[i++], "%d", &SleepSecs);
    strncpy(OutFilePrefix, argv[i++], 256);
}

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;
};

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 handleMPIerror(FILE* fpLog, const string str, int retErr, const MPI_Status* Status = NULL)
{
    char  errString[MPI_MAX_ERROR_STRING];
    int   errStrLen;
    MPI_Error_string(retErr, errString, &errStrLen);

    fprintf(fpLog, "\n\n==========================================\n");
    fprintf(fpLog, "%s: \n", str.c_str());
    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);
}

int mainInit(int argc, char* argv[], int* clusterSize, int* procRank, FILE** fppLog)
{
    if(argc < 2)
    {
        // PrintUsage(ProgName);
        PrintUsage(argv[0]);
        return -1;
    }

    // CreateOwnSignalHandler();

    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;
    char hostname[256];

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

    char filename[256];
    sprintf(filename, "./%s_r%d.log", OutFilePrefix, rank);

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

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

    *clusterSize = size;
    *procRank    = rank;

    return 0;
}

void Rcv_WaitAny(int Slaves, int iters, FILE* fpLog)
{
    Statistics myStatistics(freqSecReport, iters * Slaves);

    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>         tags(Slaves);           // Save previous recv tag, and swap it to new one on next generation
    vector<int>         SlavesRcvIters(Slaves); // count number of recv iterations per slave
    vector<MPI_Request> RcvRequests(Slaves);


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

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

        if  (retErr != MPI_SUCCESS) {
            char msg[256];
            sprintf(msg, "Rank %d, fail - request deallocated", slaveRank);
            handleMPIerror(fpLog, msg, retErr, &status);
            exit(-1);
        }

        ++SlavesRcvIters[slaveIdx]; // received one message
        myStatistics.print(fpLog);

        usleep(500); // some processing, network trafic full degradation


        // swap tag & enter blocked recv
        MPI_Status stat;
        tags[slaveIdx] = (tags[slaveIdx] == TAG1) ? TAG2 : TAG1;
        MPI_Recv(RcvBufs[slaveIdx], BUF_SZ, MPI::CHAR, slaveRank, tags[slaveIdx], MPI_COMM_WORLD, &stat);

        ++SlavesRcvIters[slaveIdx];
        myStatistics.print(fpLog);


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

        if(SlavesRcvIters[slaveIdx] == nIters)
        {
            ++SlavesFinished;
            int* testNum = (int*)(RcvBufs[slaveIdx]);
            fprintf(fpLog, "\n\nSlave finished: rank %d, last number = %d Iters = %d\n ", slaveRank, *testNum, SlavesRcvIters[slaveIdx]);
        }
        else
        {
            //swap tag and enter to unblocked recv
            tags[slaveIdx] = (tags[slaveIdx] == TAG1) ? TAG2 : TAG1;
            MPI_Irecv(RcvBufs[slaveIdx], BUF_SZ, MPI::CHAR, slaveRank, tags[slaveIdx], MPI_COMM_WORLD, &(RcvRequests[slaveIdx]));
        }
        fflush(fpLog);
    }

    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];

    int TAG = TAG1; // start sending with TAG1

    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, TAG, MPI_COMM_WORLD);
        if(retErr != MPI_SUCCESS) {
            char msg[256];
            sprintf(msg, "MPI_Send Recognized error rank - %d iter = %d, fail", rank, i);
            handleMPIerror(fpLog, msg, retErr);
            return -1;
        }

        myStatistics.print(fpLog);
        //swap tag
        TAG = (TAG == TAG1) ? TAG2 : TAG1;
    }

    fprintf(fpLog, "Last Sent number = %d\n", *(int*)chBuf);
    fprintf(fpLog, "\n\n\nEnd of slave subroutine\n");
    fflush(fpLog);
    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");
    }

    return 0;
}






