#include "pid_controller.h"

static float PID_Controller_Clamp(float value, float minimum, float maximum)
{
    if (value > maximum)
    {
        return maximum;
    }
    if (value < minimum)
    {
        return minimum;
    }
    return value;
}

void PID_Controller_Init(PID_Controller_t *pid,
                         float kp,
                         float ki,
                         float kd,
                         float sample_time,
                         float output_min,
                         float output_max,
                         float integral_limit)
{
    pid->target = 0.0f;
    pid->measurement = 0.0f;
    pid->output = 0.0f;
    pid->kp = kp;
    pid->ki = ki;
    pid->kd = kd;
    pid->sample_time = sample_time;
    pid->integral_term = 0.0f;
    pid->previous_measurement = 0.0f;
    pid->integral_limit = integral_limit;
    pid->output_min = output_min;
    pid->output_max = output_max;
    pid->conditional_integration = 1U;
}

void PID_Controller_Reset(PID_Controller_t *pid, float measurement)
{
    pid->measurement = measurement;
    pid->output = 0.0f;
    pid->integral_term = 0.0f;
    pid->previous_measurement = measurement;
}

float PID_Controller_Update(PID_Controller_t *pid,
                            float target,
                            float measurement)
{
    float error = target - measurement;
    float proportional_term = pid->kp * error;
    float derivative_term = 0.0f;
    float candidate_integral;
    float raw_output;
    uint8_t saturated_in_helpful_direction;

    pid->target = target;
    pid->measurement = measurement;

    if (pid->sample_time > 0.0f)
    {
        derivative_term = -pid->kd *
                          (measurement - pid->previous_measurement) /
                          pid->sample_time;
    }

    candidate_integral = pid->integral_term +
                         pid->ki * error * pid->sample_time;
    candidate_integral = PID_Controller_Clamp(candidate_integral,
                                              -pid->integral_limit,
                                              pid->integral_limit);

    raw_output = proportional_term + candidate_integral + derivative_term;
    saturated_in_helpful_direction =
        ((raw_output > pid->output_max) && (error > 0.0f)) ||
        ((raw_output < pid->output_min) && (error < 0.0f));

    if (!pid->conditional_integration || !saturated_in_helpful_direction)
    {
        pid->integral_term = candidate_integral;
    }

    raw_output = proportional_term + pid->integral_term + derivative_term;
    pid->output = PID_Controller_Clamp(raw_output,
                                       pid->output_min,
                                       pid->output_max);
    pid->previous_measurement = measurement;

    return pid->output;
}
