// #include "mpi_test_incl.h"
#include <mpi.h>
#include <iostream>
#include <vector>
#include <string.h>
#include <stdlib.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 = 100;

using namespace std;

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

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[])
{
    int i = 1;
    sscanf(argv[i++], "%d", &nIters);
    sscanf(argv[i++], "%d", &Scale);
    strncpy(OutFilePrefix, argv[i++], 256);
}

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)
{
    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) {
            fprintf(fpLog, "Rank %d, fail - request deallocated", slaveRank);
            exit(-1);
        }

        ++SlavesRcvIters[slaveIdx]; // received one message

//         usleep(200000); // 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];
*/

        // 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");
}


void SndSyncSlave(int rank, FILE* fpLog)
{
    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) {
            fprintf(fpLog, "Waitall Recognized error rank - %d iter = %d, fail", rank, i);
            exit(-1);
        }

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


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






