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

int main(void) {
    int rank, size;
    MPI_Init(NULL, NULL);
    MPI_Comm_rank(MPI_COMM_WORLD, &rank);
    MPI_Comm_size(MPI_COMM_WORLD, &size);

    if (size != 4) {
        if (rank == 0) printf("Run this example with four processes.\n");
        MPI_Finalize();
        return 1;
    }

    int input[4] = {8, 3, 6, 1};
    int result[4];
    int value;
    MPI_Scatter(input, 1, MPI_INT, &value, 1, MPI_INT,
                0, MPI_COMM_WORLD);

    for (int phase = 0; phase < size; phase++) {
        int partner;
        if (phase % 2 == 0) {
            partner = (rank % 2 == 0) ? rank + 1 : rank - 1;
        } else {
            partner = (rank % 2 == 0) ? rank - 1 : rank + 1;
        }

        if (partner >= 0 && partner < size) {
            int neighbor_value;
            MPI_Sendrecv(&value, 1, MPI_INT, partner, 0,
                         &neighbor_value, 1, MPI_INT, partner, 0,
                         MPI_COMM_WORLD, MPI_STATUS_IGNORE);
            if (rank < partner && value > neighbor_value) value = neighbor_value;
            if (rank > partner && value < neighbor_value) value = neighbor_value;
        }
    }

    MPI_Gather(&value, 1, MPI_INT, result, 1, MPI_INT,
               0, MPI_COMM_WORLD);
    if (rank == 0) {
        printf("Before: 8 3 6 1\n");
        printf("After:  %d %d %d %d\n",
               result[0], result[1], result[2], result[3]);
    }

    MPI_Finalize();
    return 0;
}
