/*
  Parallel Programming - 2026/2027
  Implementation of a termination detection algorithm (Dijkstra-Scholten)
*/

#include <mpi.h>
#include <pthread.h>
#include <stdio.h>
#include <stdbool.h>
#include <unistd.h>

#include "global.h"
#include "control.h"
#include "util.h"

// Tag for control messages. Distinct from basic algorithm's tags (1).
#define CONTROL_SIGNAL 99

// Message types for control routing
typedef enum {
    MSG_ACK,
    MSG_KILL
} msg_type_t;

// Type for control messages
typedef struct {
    msg_type_t type;
} control_message_t;

// Node operational states
typedef enum {
    ACTIVE_STATE,
    PASSIVE_STATE
} node_state_t;


// SHARED LOCAL STATE - Must be protected by mutex because basic_thread and control_thread access it concurrently
static pthread_mutex_t state_mutex = PTHREAD_MUTEX_INITIALIZER;
static node_state_t    state       = PASSIVE_STATE;
static int             parent      = -1; // -1 indicates null/no parent
static int             deficit     = 0;  // C_i counter (unacknowledged messages)
static bool            is_initiator = false;
static bool            started      = false; // RACE CONDITION FIX


// HELPER FUNCTION - Evaluates if the node can detach from the tree. Caller must hold state_mutex.
static void try_resolve_tree(int my_id){
    if (started && state == PASSIVE_STATE && deficit == 0){
        if(is_initiator){
            // Root is passive and deficit is 0 -> Global Termination!
            return; 
        }

        if(parent != -1){
            control_message_t sig = { MSG_ACK };
            // Send acknowledgment signal up the tree to our parent
            MPI_Send(&sig, sizeof(control_message_t), MPI_BYTE, parent, 
                     CONTROL_SIGNAL, MPI_COMM_WORLD);
            
            trace("%d: [CONTROL] Sent tree-signal to parent %d\n", my_id, parent);
            parent = -1; // Detach from tree
        }
    }
}


//  Main loop of the termination detection algorithm.
void *detect_termination(void *_args){
    thread_args_t *args = _args;
    int id = args->id;
    int processes = args->processes; //  MPI interface - needed for the KILL broadcast
    is_initiator = args->initiator;

    while(true){
        pthread_mutex_lock(&state_mutex);
        
        // CHECK GLOBAL TERMINATION CONDITION - Only evaluate if the basic algorithm has officially started
        if(started && is_initiator && state == PASSIVE_STATE && deficit == 0){
            trace("%d: [CONTROL] GLOBAL TERMINATION DETECTED!\n", id);
            
            // Broadcast KILL signal to unblock all other control threads - use 
            control_message_t kill_sig = {MSG_KILL};
            for(int p = 0; p < processes; p++){
                if(p != id){
                    MPI_Send(
                        &kill_sig,                      // pointer to the control_message_t struct.
                        sizeof(control_message_t),      // the size of the message.
                        MPI_BYTE,                       // sending raw memory bytes.
                        p,                              // destinnationn, parent node.
                        CONTROL_SIGNAL,                 // label 99.
                        MPI_COMM_WORLD                  // the overall group of processes this message is restricted to.
                    );
                }
            }
            
            pthread_mutex_unlock(&state_mutex);
            break; // Breaks loop, tells main() to shut down
        }
        
        pthread_mutex_unlock(&state_mutex);

        // Check network for incoming child signals (non-blocking) - use MPI_Recv to read messages sent over the network by other processes
        int flag = 0;
        MPI_Status status;
        MPI_Iprobe(MPI_ANY_SOURCE, CONTROL_SIGNAL, MPI_COMM_WORLD, &flag, &status);

        if(flag){
            control_message_t sig;
            MPI_Recv(
                &sig,                                   // pointer to the memory space where the incoming message will be stored
                sizeof(control_message_t),              // the maximum size (in bytes) of the data willing to receive.
                MPI_BYTE,                               // raw stream of bytes
                status.MPI_SOURCE,                      // the ID (rank) of the process you want to receive from MPI_Iprobe.
                CONTROL_SIGNAL,                         // label of the message
                MPI_COMM_WORLD,                         // the overall group of processes this message belongs to.
                MPI_STATUS_IGNORE                       // just to save memory and processing time.
            );

            if(sig.type == MSG_KILL){
                trace("%d: [CONTROL] Received KILL signal. Shutting down.\n", id);
                break;
            }
            else if(sig.type == MSG_ACK){
                pthread_mutex_lock(&state_mutex);
                deficit--; // Rule D: A child finished its work and reported back
                trace("%d: [CONTROL] Got signal from %d (New Deficit: %d)\n", 
                      id, status.MPI_SOURCE, deficit);
                
                try_resolve_tree(id);
                pthread_mutex_unlock(&state_mutex);
            }
        }
        else{
            // Sleep briefly to prevent pinning the CPU core at 100% usage
            usleep(1000); 
        }
    }

    return NULL;
}

/*
  Called at startup of the basic algorithm in process ID.
*/
void control_basic_start_hook(int id)
{
    pthread_mutex_lock(&state_mutex);
    is_initiator = (id == 0);
    state   = is_initiator ? ACTIVE_STATE : PASSIVE_STATE;
    parent  = -1;
    deficit = 0;
    started = true; // Signal the control thread that it is safe to check for termination
    pthread_mutex_unlock(&state_mutex);
}

/*
  Called when the basic algorithm process ID becomes passive.
*/
void control_become_passive_hook(int id)
{
    pthread_mutex_lock(&state_mutex);
    state = PASSIVE_STATE;
    try_resolve_tree(id); // Check if we can leave the tree right now
    pthread_mutex_unlock(&state_mutex);
}

/*
  Called when the basic algorithm process ID becomes active.
*/
void control_become_active_hook(int id)
{
    pthread_mutex_lock(&state_mutex);
    state = ACTIVE_STATE;
    pthread_mutex_unlock(&state_mutex);
}

/*
  Called when the basic algorithm process ID sends a basic message to
  process PEER.
*/
void control_basic_send_hook(int id, int peer)
{
    pthread_mutex_lock(&state_mutex);
    deficit++; // Rule A: Sent unacknowledged work
    pthread_mutex_unlock(&state_mutex);
}

/*
  Called when the basic algorithm process ID receives a basic message
  from process PEER.
*/
void control_basic_receive_hook(int id, int peer)
{
    pthread_mutex_lock(&state_mutex);
    
    // A node is only "out of the tree" if parent == -1.
    // It can be PASSIVE but still in the tree if it is waiting for children (deficit > 0).
    if(parent == -1 && !is_initiator){
        // Not in the tree -> Join it
        parent = peer; 
        trace("%d: [CONTROL] Joined tree under parent %d\n", id, parent);
    }else{
        // Already in the tree (or is root) -> Reject parent change, immediately signal back
        control_message_t sig = { MSG_ACK };
        MPI_Send(&sig, sizeof(control_message_t), MPI_BYTE, peer, 
                 CONTROL_SIGNAL, MPI_COMM_WORLD);
        trace("%d: [CONTROL] Rejected parent %d (Already in tree or Root)\n", id, peer);
    }
    
    pthread_mutex_unlock(&state_mutex);
}