#include <stdint.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>

// Storage for training images (60k × 784 bytes) – on PC we can malloc
static uint8_t *train_images = NULL;
static uint8_t *train_labels = NULL;
static uint8_t *test_images  = NULL;
static uint8_t *test_labels  = NULL;
static int train_loaded = 0;
static int test_loaded  = 0;

// Helper: read 32‑bit big‑endian integer
static int read_int(FILE *f) {
    unsigned char buf[4];
    if (fread(buf, 1, 4, f) != 4) return -1;
    return (buf[0] << 24) | (buf[1] << 16) | (buf[2] << 8) | buf[3];
}

// Load MNIST data from files (call once before using get_* functions)
void load_mnist_data(const char *train_img_path, const char *train_lbl_path,
                     const char *test_img_path,  const char *test_lbl_path) {
    FILE *f;
    int magic, num_items, rows, cols;

    // Training images
    f = fopen(train_img_path, "rb");
    if (!f) { perror("train images"); exit(1); }
    magic = read_int(f);
    num_items = read_int(f);
    rows = read_int(f);
    cols = read_int(f);
    if (magic != 0x803 || rows != 28 || cols != 28) {
        fprintf(stderr, "Invalid training image file\n");
        exit(1);
    }
    train_images = malloc(num_items * rows * cols);
    fread(train_images, 1, num_items * rows * cols, f);
    fclose(f);
    train_loaded = num_items;

    // Training labels
    f = fopen(train_lbl_path, "rb");
    if (!f) { perror("train labels"); exit(1); }
    magic = read_int(f);
    num_items = read_int(f);
    if (magic != 0x801) { fprintf(stderr, "Invalid training label file\n"); exit(1); }
    train_labels = malloc(num_items);
    fread(train_labels, 1, num_items, f);
    fclose(f);

    // Test images
    f = fopen(test_img_path, "rb");
    if (!f) { perror("test images"); exit(1); }
    magic = read_int(f);
    num_items = read_int(f);
    rows = read_int(f);
    cols = read_int(f);
    if (magic != 0x803 || rows != 28 || cols != 28) {
        fprintf(stderr, "Invalid test image file\n");
        exit(1);
    }
    test_images = malloc(num_items * rows * cols);
    fread(test_images, 1, num_items * rows * cols, f);
    fclose(f);
    test_loaded = num_items;

    // Test labels
    f = fopen(test_lbl_path, "rb");
    if (!f) { perror("test labels"); exit(1); }
    magic = read_int(f);
    num_items = read_int(f);
    if (magic != 0x801) { fprintf(stderr, "Invalid test label file\n"); exit(1); }
    test_labels = malloc(num_items);
    fread(test_labels, 1, num_items, f);
    fclose(f);
}

// External functions called from CCT‑Lang
void get_train_image(int idx, uint8_t out[784]) {
    if (idx < 0 || idx >= train_loaded) return;
    memcpy(out, train_images + idx * 784, 784);
}

uint8_t get_train_label(int idx) {
    if (idx < 0 || idx >= train_loaded) return 0;
    return train_labels[idx];
}

void get_test_image(int idx, uint8_t out[784]) {
    if (idx < 0 || idx >= test_loaded) return;
    memcpy(out, test_images + idx * 784, 784);
}

uint8_t get_test_label(int idx) {
    if (idx < 0 || idx >= test_loaded) return 0;
    return test_labels[idx];
}