#include <math.h>
#include <CircularBuffer.hpp>
#include <BLEDevice.h>
#include <BLEServer.h>
#include <BLEUtils.h>
#include <BLE2902.h>

#define SAMPLE_RATE 125
#define INPUT_PIN A0
#define DATA_LENGTH 16
#define RECORD_DURATION_MS 30000

enum Mode { STANDBY, RECORDING };
Mode currentMode = STANDBY;
unsigned long recordingStartTime = 0;
int ts = 0;

// BLE UUIDs
#define SERVICE_UUID       "f77a2094-8d48-4c1b-bb58-f6a794c63e80"
#define CHAR_CONTROL_UUID  "cf84636d-2f54-47d4-a665-09c06a252fea"
#define CHAR_DATA_UUID     "38730795-f977-44b1-8d1d-baa9cc370522"
#define CHAR_REPORT_UUID   "870df317-e2e8-4010-b596-2d175092a3fe"

BLECharacteristic* dataChar;
BLECharacteristic* reportChar;

bool peak = false;
bool IgnoreReading = false;
bool FirstPulseDetected = false;
unsigned long FirstPulseTime = 0;
unsigned long rrInterval = 0;
unsigned long systoleStart = 0;
unsigned long diastoleStart = 0;
float rmssd = 0.0;
bool isSystole = false;
bool isDiastole = false;

CircularBuffer<unsigned long, 10> rrBuffer;
CircularBuffer<unsigned long, 10> rrDiffSqBuffer;
CircularBuffer<int, 1000> bpmSamples;
CircularBuffer<unsigned long, 1000> rrSamples;
CircularBuffer<float, 1000> rmssdSamples;
CircularBuffer<unsigned long, 1000> systoleSamples;
CircularBuffer<unsigned long, 1000> diastoleSamples;
int data_index = 0;

// BLE server pointer for advertising
BLEServer* pServer;

// --- Callbacks ---

class ControlCallback : public BLECharacteristicCallbacks {
    void onWrite(BLECharacteristic* pCharacteristic) override {
        String value = pCharacteristic->getValue();
        if (value == "start" && currentMode == STANDBY) {
            currentMode = RECORDING;
            recordingStartTime = millis();
            Serial.println("BLE received 'start'. Recording begun.");
        }
    }
};

class ServerCallbacks : public BLEServerCallbacks {
    void onConnect(BLEServer* pServer) override {
        Serial.println("Client connected");
    }
    void onDisconnect(BLEServer* pServer) override {
        Serial.println("Client disconnected. Restarting advertising...");
        pServer->getAdvertising()->start();
    }
};

void sendReport();
void resetBuffers();

void setup() {
    Serial.begin(115200);
    delay(300); // Optional for Serial stabilization

    pinMode(INPUT_PIN, INPUT);

    BLEDevice::init("Cardiogram");
    pServer = BLEDevice::createServer();
    pServer->setCallbacks(new ServerCallbacks());

    BLEService* pService = pServer->createService(SERVICE_UUID);

    BLECharacteristic* controlChar = pService->createCharacteristic(
        CHAR_CONTROL_UUID, BLECharacteristic::PROPERTY_WRITE
    );
    controlChar->setCallbacks(new ControlCallback());
    controlChar->addDescriptor(new BLE2902());

    dataChar = pService->createCharacteristic(
        CHAR_DATA_UUID, BLECharacteristic::PROPERTY_NOTIFY
    );
    dataChar->addDescriptor(new BLE2902());

    reportChar = pService->createCharacteristic(
        CHAR_REPORT_UUID, BLECharacteristic::PROPERTY_NOTIFY
    );
    reportChar->addDescriptor(new BLE2902());

    pService->start();
    pServer->getAdvertising()->start();

    Serial.println("Setup complete, advertising started.");
}

void loop() {
    if (currentMode == STANDBY) {
        delay(10);
        ts = 0;
        return;
    }

    static unsigned long past = 0;
    unsigned long present = micros();
    unsigned long interval = present - past;
    past = present;
    static long timer = 0;
    timer -= interval;

    if (timer < 0) {
        timer += 1000000 / SAMPLE_RATE;

        int adcValue = analogRead(INPUT_PIN);
        float signal = ECGFilter(adcValue) / 512.0;
        peak = GetPeak(signal);

        if (peak && !IgnoreReading) {
            unsigned long now = millis();
            if (FirstPulseDetected) {
                rrInterval = now - FirstPulseTime;
                rrBuffer.unshift(rrInterval);
                rrSamples.push(rrInterval);
                unsigned long systoleDuration = now - systoleStart;
                unsigned long diastoleDuration = systoleStart - diastoleStart;
                systoleSamples.push(systoleDuration);
                diastoleSamples.push(diastoleDuration);
                isSystole = true;
                isDiastole = false;

                if (rrBuffer.size() >= 2) {
                    long diff = (long)rrBuffer[0] - (long)rrBuffer[1];
                    rrDiffSqBuffer.unshift(diff * diff);
                    float sum = 0.0;
                    for (int i = 0; i < rrDiffSqBuffer.size(); i++) sum += rrDiffSqBuffer[i];
                    rmssd = sqrt(sum / rrDiffSqBuffer.size());
                    rmssdSamples.push(rmssd);
                }
                systoleStart = now;
            } else {
                systoleStart = now;
                diastoleStart = now;
                FirstPulseDetected = true;
            }
            FirstPulseTime = now;
            IgnoreReading = true;
        } else if (!peak) {
            IgnoreReading = false;
            isSystole = false;
            isDiastole = true;
            if (diastoleStart == 0) diastoleStart = millis();
        }

        int currentBPM = 0;
        if (rrBuffer.isFull()) {
            uint32_t avgRR = 0;
            for (int i = 0; i < rrBuffer.size(); i++) avgRR += rrBuffer[i];
            avgRR /= rrBuffer.size();
            currentBPM = (1.0 / avgRR) * 60000;
            bpmSamples.push(currentBPM);
        }

        // Prepare metric string for both Serial and BLE
        char buf[512];
        snprintf(buf, sizeof(buf), "{ \"status\": \"STREAMING\", \"adcValue\": %d,\"currentBPM\": %d, \"rrInterval\": %lu, \"rmssd\": %.2f, \"avgSystole\": %lu, \"avgDiastole\": %lu, \"isSystole\": %d, \"isDiastole\": %d, \"index\": %d }", adcValue, currentBPM, rrInterval, rmssd,
                 systoleSamples[0], diastoleSamples[0], isSystole, isDiastole, ts++);

        // Show on Serial
        Serial.println(buf);
        
        // Send to BLE
        dataChar->setValue(buf);
        dataChar->notify();

        // Auto-report every RECORD_DURATION_MS
        if (millis() - recordingStartTime >= RECORD_DURATION_MS) {
            delay(150);
            sendReport();
            resetBuffers();
            currentMode = STANDBY;
        }
    }

    delay(10);
}

void sendReport() {
    char report[512];
    snprintf(report, sizeof(report),
             "{\"status\": \"FINISHED\", \"avgBPM\":%d, \"avgRR\":%lu, \"avgRMSSD\":%.2f, \"avgSystole\":%lu, \"avgDiastole\":%lu}",
             averageIntBuffer(bpmSamples),
             averageLongBuffer(rrSamples),
             averageFloatBuffer(rmssdSamples),
             averageLongBuffer(systoleSamples),
             averageLongBuffer(diastoleSamples));
    // Print to Serial for debug
    Serial.println(report);

    // Send via BLE
    reportChar->setValue(report);
    reportChar->notify();
    reportChar->notify();
    reportChar->notify();
}

void resetBuffers() {
    bpmSamples.clear();
    rrSamples.clear();
    rmssdSamples.clear();
    systoleSamples.clear();
    diastoleSamples.clear();
    rrBuffer.clear();
    rrDiffSqBuffer.clear();
    FirstPulseDetected = false;
    IgnoreReading = false;
    ts = 0;
}

int averageIntBuffer(CircularBuffer<int, 1000> &buf) {
    if (buf.size() == 0) return 0;
    long sum = 0;
    for (int i = 0; i < buf.size(); i++) sum += buf[i];
    return sum / buf.size();
}

unsigned long averageLongBuffer(CircularBuffer<unsigned long, 1000> &buf) {
    if (buf.size() == 0) return 0;
    unsigned long sum = 0;
    for (int i = 0; i < buf.size(); i++) sum += buf[i];
    return sum / buf.size();
}

float averageFloatBuffer(CircularBuffer<float, 1000> &buf) {
    if (buf.size() == 0) return 0.0;
    float sum = 0;
    for (int i = 0; i < buf.size(); i++) sum += buf[i];
    return sum / buf.size();
}

bool GetPeak(float new_sample) {
    static float data_buffer[DATA_LENGTH];
    static float mean_buffer[DATA_LENGTH];
    static float stddev_buffer[DATA_LENGTH];
    bool peak_detected = false;

    if (new_sample - mean_buffer[data_index] > (DATA_LENGTH / 2.0) * stddev_buffer[data_index]) {
        data_buffer[data_index] = new_sample + data_buffer[data_index];
        peak_detected = true;
    } else {
        data_buffer[data_index] = new_sample;
    }
    float sum = 0.0, mean = 0.0, stddev = 0.0;
    for (int i = 0; i < DATA_LENGTH; ++i) sum += data_buffer[(data_index + i) % DATA_LENGTH];
    mean = sum / DATA_LENGTH;
    for (int i = 0; i < DATA_LENGTH; ++i)
        stddev += pow(data_buffer[(i) % DATA_LENGTH] - mean, 2);
    mean_buffer[data_index] = mean;
    stddev_buffer[data_index] = sqrt(stddev / DATA_LENGTH);
    data_index = (data_index + 1) % DATA_LENGTH;
    return peak_detected;
}

float ECGFilter(float input) {
    float output = input;
    {
        static float z1, z2;
        float x = output - 0.70682283 * z1 - 0.15621030 * z2;
        output = 0.28064917 * x + 0.56129834 * z1 + 0.28064917 * z2;
        z2 = z1; z1 = x;
    }
    {
        static float z1, z2;
        float x = output - 0.95028224 * z1 - 0.54073140 * z2;
        output = 1.0 * x + 2.0 * z1 + 1.0 * z2;
        z2 = z1; z1 = x;
    }
    {
        static float z1, z2;
        float x = output - -1.95360385 * z1 - 0.95423412 * z2;
        output = 1.0 * x - 2.0 * z1 + 1.0 * z2;
        z2 = z1; z1 = x;
    }
    {
        static float z1, z2;
        float x = output - -1.98048558 * z1 - 0.98111344 * z2;
        output = 1.0 * x - 2.0 * z1 + 1.0 * z2;
        z2 = z1; z1 = x;
    }
    return output;
}