#include <Arduino.h>
#include <Encoder.h>
#include <Servo.h>

// Encoder setup
Encoder myEnc(20, 21);
const int PPR = 600;
const int countsPerRev = PPR * 4; // 2400 for quadrature

// Wheel & motion constants
const float wheelDiameter = 0.0508;
const float wheelCircumference = PI * wheelDiameter;

// Control pins
const int escPin = 9;
const int buttonPin = 12;
Servo motor;

// Trapezoid motion profile
float Distance = 8;
float targetDistance = Distance - .7;
float accelDistance = 1.25;
float decelDistance = 4.0;
float maxVelocity = 3.0; // meters per second

// PID control
float Kp = 3000; // Tune these
float Ki = 0;
float Kd = 100;

float desiredVelocity = 0;
float actualVelocity = 0;
float lastPosition = 0;
float lastError = 0;
float integral = 0;

// ESC pulse range
const int neutralPWM = 1500;
const int minPWM = 1550;
const int maxPWM = 1600;

// State
bool isMoving = false;
unsigned long lastUpdateTime = 0;
unsigned long startTime = 0;

void setup() {
  motor.attach(escPin);
  motor.writeMicroseconds(neutralPWM);
  pinMode(buttonPin, INPUT_PULLUP);
  Serial.begin(9600);
  Serial.println("Ready");
}

void loop() {
  long encoderCounts = myEnc.read();
  float currentPosition = (encoderCounts / (float)countsPerRev) * wheelCircumference;

  unsigned long now = millis();
  float dt = (now - lastUpdateTime) / 1000.0;

  // Start motion
  if (digitalRead(buttonPin) == LOW && !isMoving) {
    isMoving = true;
    myEnc.write(0);
    lastPosition = 0;
    lastUpdateTime = now;
    startTime = now;
    Serial.println("Motion started");
  }

  if (!isMoving) return;

  // Calculate actual velocity
  if (dt > 0.01) {
    actualVelocity = (currentPosition - lastPosition) / dt;
    lastPosition = currentPosition;
    lastUpdateTime = now;
  }

  // Compute distance remaining
  float distanceRemaining = targetDistance - currentPosition;

  // Check if motion is complete
  if (distanceRemaining <= 0) {
    motor.writeMicroseconds(neutralPWM);
    isMoving = false;
    Serial.println("Target reached.");
    return;
  }

  // Trapezoidal motion profile
  if (currentPosition < accelDistance) {
    float factor = currentPosition / accelDistance;
    desiredVelocity = maxVelocity * constrain(factor, 0.0, 1.0);
  } else if (currentPosition >= accelDistance && currentPosition < (targetDistance - decelDistance)) {
    desiredVelocity = maxVelocity;
  } else {
    float factor = distanceRemaining / decelDistance;
    desiredVelocity = maxVelocity * constrain(factor, 0.0, 1.0);
  }

  // PID velocity control
  float error = desiredVelocity - actualVelocity;
  integral += error * dt;
  float derivative = (error - lastError) / dt;
  float output = Kp * error + Ki * integral + Kd * derivative;
  lastError = error;

  // Convert PID output to PWM signal
  int pwmSignal = neutralPWM + output;
  pwmSignal = constrain(pwmSignal, minPWM, maxPWM);
  motor.writeMicroseconds(pwmSignal);
  //pio run -e mega_comp -t upload

  // Debug output
  Serial.print("Pos: "); Serial.print(currentPosition, 3);
  Serial.print(" m, Vset: "); Serial.print(desiredVelocity, 2);
  Serial.print(" m/s, Vact: "); Serial.print(actualVelocity, 2);
  Serial.print(" m/s, PWM: "); Serial.println(pwmSignal);
}
