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

int main(int argc, char *argv[]) {
    int rank, size;
    MPI_Init(NULL, NULL);
    MPI_Comm_rank(MPI_COMM_WORLD, &rank);
    MPI_Comm_size(MPI_COMM_WORLD, &size);

    long long total_tosses = 10000000;
    if (rank == 0) {
        if (argc > 2) {
            total_tosses = 0;
        } else if (argc == 2) {
            total_tosses = atoll(argv[1]);
        }
    }
    MPI_Bcast(&total_tosses, 1, MPI_LONG_LONG_INT, 0, MPI_COMM_WORLD);
    if (total_tosses < 1 || total_tosses > 1000000000LL) {
        if (rank == 0) printf("Use: ./mpi_pi [tosses from 1 to 1000000000]\n");
        MPI_Finalize();
        return 1;
    }

    long long local_tosses = total_tosses / size;
    if (rank < total_tosses % size) local_tosses++;
    srand(12345u + (unsigned int) rank);
    long long local_hits = 0, total_hits = 0;

    MPI_Barrier(MPI_COMM_WORLD);
    double start = MPI_Wtime();
    for (long long toss = 0; toss < local_tosses; toss++) {
        double x = 2.0 * rand() / (RAND_MAX + 1.0) - 1.0;
        double y = 2.0 * rand() / (RAND_MAX + 1.0) - 1.0;
        if (x * x + y * y <= 1.0) local_hits++;
    }
    MPI_Reduce(&local_hits, &total_hits, 1, MPI_LONG_LONG_INT,
               MPI_SUM, 0, MPI_COMM_WORLD);
    double elapsed = MPI_Wtime() - start;
    double max_elapsed = 0.0;
    MPI_Reduce(&elapsed, &max_elapsed, 1, MPI_DOUBLE,
               MPI_MAX, 0, MPI_COMM_WORLD);

    if (rank == 0) {
        double estimate = 4.0 * total_hits / total_tosses;
        double error = estimate - 3.141592653589793;
        if (error < 0) error = -error;
        printf("Tosses: %lld; processes: %d\n", total_tosses, size);
        printf("Hits: %lld\n", total_hits);
        printf("Pi estimate: %.8f; absolute error: %.8f\n", estimate, error);
        printf("Timed calculation: %.6f seconds (maximum rank time)\n", max_elapsed);
    }
    MPI_Finalize();
    return 0;
}
