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

#define N 12

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 (N % size != 0) {
        if (rank == 0) printf("Choose a process count that divides %d.\n", N);
        MPI_Finalize();
        return 1;
    }

    int local_n = N / size;
    double start[N] = {0};
    double finish[N] = {0};
    double local[N] = {0};
    double next[N] = {0};
    double left_edge = 0.0, right_edge = 0.0;

    if (rank == 0) start[N / 2] = 100.0;
    MPI_Scatter(start, local_n, MPI_DOUBLE,
                local, local_n, MPI_DOUBLE, 0, MPI_COMM_WORLD);

    int left_rank = rank == 0 ? MPI_PROC_NULL : rank - 1;
    int right_rank = rank == size - 1 ? MPI_PROC_NULL : rank + 1;

    /* First values travel left; last values travel right. */
    MPI_Sendrecv(&local[0], 1, MPI_DOUBLE, left_rank, 0,
                 &right_edge, 1, MPI_DOUBLE, right_rank, 0,
                 MPI_COMM_WORLD, MPI_STATUS_IGNORE);
    MPI_Sendrecv(&local[local_n - 1], 1, MPI_DOUBLE, right_rank, 1,
                 &left_edge, 1, MPI_DOUBLE, left_rank, 1,
                 MPI_COMM_WORLD, MPI_STATUS_IGNORE);

    for (int i = 0; i < local_n; i++) {
        int global_i = rank * local_n + i;
        if (global_i == 0 || global_i == N - 1) {
            next[i] = 0.0;
        } else {
            double left = i == 0 ? left_edge : local[i - 1];
            double right = i == local_n - 1 ? right_edge : local[i + 1];
            next[i] = (left + 2.0 * local[i] + right) / 4.0;
        }
    }

    MPI_Gather(next, local_n, MPI_DOUBLE,
               finish, local_n, MPI_DOUBLE, 0, MPI_COMM_WORLD);

    if (rank == 0) {
        printf("After one step:");
        for (int i = 0; i < N; i++) printf(" %.0f", finish[i]);
        printf("\n");
    }

    MPI_Finalize();
    return 0;
}
