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

#include <sodium.h>

#define DEFAULT_LENGTH 20
#define MIN_LENGTH 4
#define MAX_LENGTH 128

static const char *GROUPS[] = {
    "abcdefghijklmnopqrstuvwxyz",
    "ABCDEFGHIJKLMNOPQRSTUVWXYZ",
    "0123456789",
    "!@#$%^&*()-_=+"
};

static const char ALPHABET[] =
    "abcdefghijklmnopqrstuvwxyz"
    "ABCDEFGHIJKLMNOPQRSTUVWXYZ"
    "0123456789"
    "!@#$%^&*()-_=+";

static bool contains_group(const char *password, size_t length,
                           const char *group) {
    for (size_t index = 0; index < length; index++) {
        if (strchr(group, password[index]) != NULL) {
            return true;
        }
    }
    return false;
}

static bool satisfies_policy(const char *password, size_t length) {
    const size_t group_count = sizeof GROUPS / sizeof GROUPS[0];

    for (size_t group = 0; group < group_count; group++) {
        if (!contains_group(password, length, GROUPS[group])) {
            return false;
        }
    }
    return true;
}

static int generate_password(char *output, size_t length) {
    if (output == NULL || length < MIN_LENGTH || length > MAX_LENGTH) {
        return -1;
    }

    const uint32_t alphabet_size = (uint32_t)(sizeof ALPHABET - 1);

    do {
        for (size_t index = 0; index < length; index++) {
            output[index] = ALPHABET[randombytes_uniform(alphabet_size)];
        }
    } while (!satisfies_policy(output, length));

    output[length] = '\0';
    return 0;
}

static int parse_length(const char *text, size_t *length) {
    char *end = NULL;
    errno = 0;
    const long value = strtol(text, &end, 10);

    if (errno != 0 || end == text || *end != '\0' ||
        value < MIN_LENGTH || value > MAX_LENGTH) {
        return -1;
    }

    *length = (size_t)value;
    return 0;
}

int main(int argc, char **argv) {
    size_t length = DEFAULT_LENGTH;

    if (argc > 2 ||
        (argc == 2 && parse_length(argv[1], &length) != 0)) {
        fprintf(stderr, "Usage: %s [length from %d to %d]\n",
                argv[0], MIN_LENGTH, MAX_LENGTH);
        return EXIT_FAILURE;
    }

    if (sodium_init() < 0) {
        fputs("Unable to initialize libsodium.\n", stderr);
        return EXIT_FAILURE;
    }

    char password[MAX_LENGTH + 1];
    if (generate_password(password, length) != 0) {
        fputs("Unable to generate the password.\n", stderr);
        return EXIT_FAILURE;
    }

    puts(password);
    sodium_memzero(password, sizeof password);
    return EXIT_SUCCESS;
}
