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

#define N 12

int main(void) {
    int my_rank, comm_sz;
    MPI_Init(NULL, NULL);
    MPI_Comm_rank(MPI_COMM_WORLD, &my_rank);
    MPI_Comm_size(MPI_COMM_WORLD, &comm_sz);

    if (N % comm_sz != 0) {
        if (my_rank == 0) {
            printf("Choose a process count that divides %d.\n", N);
        }
        MPI_Finalize();
        return 1;
    }

    int local_n = N / comm_sz;
    int numbers[N];
    int squares[N];
    int local_numbers[N];
    int local_squares[N];

    if (my_rank == 0) {
        for (int i = 0; i < N; i++) {
            numbers[i] = i + 1;
        }
    }

    MPI_Scatter(numbers, local_n, MPI_INT,
                local_numbers, local_n, MPI_INT,
                0, MPI_COMM_WORLD);

    for (int i = 0; i < local_n; i++) {
        local_squares[i] = local_numbers[i] * local_numbers[i];
    }

    MPI_Gather(local_squares, local_n, MPI_INT,
               squares, local_n, MPI_INT,
               0, MPI_COMM_WORLD);

    if (my_rank == 0) {
        printf("Squares:");
        for (int i = 0; i < N; i++) {
            printf(" %d", squares[i]);
        }
        printf("\n");
    }

    MPI_Finalize();
    return 0;
}
