#include "PID.h"

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

void PID_Init(PID_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->actual = 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_error = 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_Reset(PID_t *pid, float measurement)
{
    pid->actual = measurement;
    pid->output = 0.0f;
    pid->integral_term = 0.0f;
    pid->previous_error = 0.0f;
    pid->previous_measurement = measurement;
}

float PID_Update(PID_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->actual = measurement;

    if (pid->sample_time > 0.0f)
    {
        /* Differentiate the measurement to avoid a setpoint kick. */
        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_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_Clamp(raw_output, pid->output_min, pid->output_max);
    pid->previous_error = error;
    pid->previous_measurement = measurement;

    return pid->output;
}
