// SPDX-License-Identifier: GPL-3.0-or-later
//
// ObsidianDart: encrypts and decrypts a local credentials file with
// AES-256-GCM and a PBKDF2-HMAC-SHA256 derived key, via OpenSSL.
// Copyright (C) 2026 Sanchez Performance LLC
//
// This program is free software: you can redistribute it and/or modify
// it under the terms of the GNU General Public License as published by
// the Free Software Foundation, either version 3 of the License, or
// (at your option) any later version.
//
// This program is distributed in the hope that it will be useful,
// but WITHOUT ANY WARRANTY; without even the implied warranty of
// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
// GNU General Public License for more details.
//
// You should have received a copy of the GNU General Public License
// along with this program.  If not, see <https://www.gnu.org/licenses/>.

#include <iostream>
#include <string>
#include <vector>
#include <cstdint>
#include <cstring>
#include <cerrno>
#include <csignal>
#include <cstdlib>
#include <openssl/evp.h>
#include <openssl/err.h>
#include <openssl/rand.h>
#include <openssl/crypto.h>
#include <fcntl.h>
#include <sys/stat.h>
#include <termios.h>
#include <unistd.h>

// File format, version 1 (all fields raw bytes, no padding):
//   magic   "OBSDART" + 0x00   8 bytes
//   version                    1 byte  (0x01)
//   salt                      16 bytes (PBKDF2 salt)
//   nonce                     12 bytes (GCM nonce)
//   ciphertext                 n bytes
//   tag                       16 bytes (GCM authentication tag)
// The header (magic, version, salt, nonce) is bound to the tag as AAD,
// so changing any byte of the file makes decryption fail.

static const char* const PLAIN_FILE = "passwords.wf";
static const char* const ENC_FILE = "passwords.wf.enc";

static const unsigned char MAGIC[8] = {'O', 'B', 'S', 'D', 'A', 'R', 'T', 0};
static const unsigned char FORMAT_VERSION = 1;
static const size_t SALT_LEN = 16;
static const size_t NONCE_LEN = 12;
static const size_t TAG_LEN = 16;
static const size_t KEY_LEN = 32;
static const int PBKDF2_ITERATIONS = 600000;
static const size_t HEADER_LEN = sizeof(MAGIC) + 1 + SALT_LEN + NONCE_LEN;

// Wipes a byte buffer before it goes out of scope.
static void wipe(std::vector<unsigned char>& v) {
    if (!v.empty()) OPENSSL_cleanse(v.data(), v.size());
    v.clear();
}

static void wipe(std::string& s) {
    if (!s.empty()) OPENSSL_cleanse(&s[0], s.size());
    s.clear();
}

static void printOpenSSLErrors() {
    ERR_print_errors_fp(stderr);
}

// Terminal state saved while echo is off, so Ctrl+C can put it back.
static struct termios savedTerm;
static volatile sig_atomic_t echoOff = 0;

static void restoreTermOnSignal(int sig) {
    if (echoOff) tcsetattr(STDIN_FILENO, TCSANOW, &savedTerm);
    const char nl = '\n';
    (void)!write(STDOUT_FILENO, &nl, 1);
    signal(sig, SIG_DFL);
    raise(sig);
}

// Colors only on a real terminal, and never when NO_COLOR is set.
static bool useColor() {
    static const int on = isatty(STDOUT_FILENO) && isatty(STDERR_FILENO) &&
                          std::getenv("NO_COLOR") == nullptr;
    return on;
}

static const char* paint(const char* code) { return useColor() ? code : ""; }
static const char* BOLD() { return paint("\033[1m"); }
static const char* DIM() { return paint("\033[2m"); }
static const char* RED() { return paint("\033[31m"); }
static const char* GREEN() { return paint("\033[32m"); }
static const char* YELLOW() { return paint("\033[33m"); }
static const char* CYAN() { return paint("\033[36m"); }
static const char* RESET() { return paint("\033[0m"); }

static void fail(const std::string& msg) {
    std::cerr << RED() << "error: " << RESET() << msg << std::endl;
}

// Reads one line from stdin with terminal echo turned off (when stdin is a
// terminal). Echo is always restored before returning, and on Ctrl+C.
static bool readSecret(const std::string& prompt, std::string& out) {
    std::cout << prompt << std::flush;
    bool isTty = (tcgetattr(STDIN_FILENO, &savedTerm) == 0);
    if (isTty) {
        struct termios newFlags = savedTerm;
        newFlags.c_lflag &= ~ECHO;
        echoOff = 1;
        tcsetattr(STDIN_FILENO, TCSANOW, &newFlags);
    }
    bool ok = static_cast<bool>(std::getline(std::cin, out));
    if (isTty) {
        tcsetattr(STDIN_FILENO, TCSANOW, &savedTerm);
        echoOff = 0;
    }
    std::cout << std::endl;
    return ok;
}

// Reads a whole file into memory. Returns false on any error.
static bool readFile(const char* path, std::vector<unsigned char>& out) {
    int fd = open(path, O_RDONLY);
    if (fd < 0) {
        std::cerr << "Error opening " << path << " for reading: " << std::strerror(errno) << std::endl;
        return false;
    }
    out.clear();
    unsigned char buf[4096];
    for (;;) {
        ssize_t n = read(fd, buf, sizeof(buf));
        if (n < 0) {
            if (errno == EINTR) continue;
            std::cerr << "Error reading " << path << ": " << std::strerror(errno) << std::endl;
            OPENSSL_cleanse(buf, sizeof(buf));
            close(fd);
            wipe(out);
            return false;
        }
        if (n == 0) break;
        out.insert(out.end(), buf, buf + n);
    }
    OPENSSL_cleanse(buf, sizeof(buf));
    close(fd);
    return true;
}

static bool writeAll(int fd, const unsigned char* data, size_t len) {
    while (len > 0) {
        ssize_t n = write(fd, data, len);
        if (n < 0) {
            if (errno == EINTR) continue;
            return false;
        }
        data += n;
        len -= static_cast<size_t>(n);
    }
    return true;
}

// Writes data to path through a temporary file (mode 0600) and an atomic
// rename, so a crash never leaves a half-written file under the real name.
static bool writeFileAtomic(const char* path, const std::vector<unsigned char>& data) {
    std::string tmp = std::string(path) + ".tmp";
    int fd = open(tmp.c_str(), O_WRONLY | O_CREAT | O_EXCL, 0600);
    if (fd < 0) {
        std::cerr << "Error creating " << tmp << ": " << std::strerror(errno) << std::endl;
        return false;
    }
    bool ok = writeAll(fd, data.data(), data.size()) && fsync(fd) == 0;
    if (close(fd) != 0) ok = false;
    if (!ok || rename(tmp.c_str(), path) != 0) {
        std::cerr << "Error writing " << path << ": " << std::strerror(errno) << std::endl;
        unlink(tmp.c_str());
        return false;
    }
    return true;
}

// Overwrites a file with zeros, syncs it, then deletes it. Best effort only:
// on SSDs and copy-on-write filesystems the old blocks may survive.
static bool shredFile(const char* path) {
    int fd = open(path, O_WRONLY);
    if (fd < 0) {
        std::cerr << "Error opening " << path << " to overwrite: " << std::strerror(errno) << std::endl;
        return false;
    }
    struct stat st;
    bool ok = (fstat(fd, &st) == 0);
    if (ok) {
        unsigned char zeros[4096];
        std::memset(zeros, 0, sizeof(zeros));
        off_t remaining = st.st_size;
        while (ok && remaining > 0) {
            size_t chunk = remaining > static_cast<off_t>(sizeof(zeros))
                               ? sizeof(zeros) : static_cast<size_t>(remaining);
            ok = writeAll(fd, zeros, chunk);
            remaining -= static_cast<off_t>(chunk);
        }
        if (ok) ok = (fsync(fd) == 0);
    }
    if (close(fd) != 0) ok = false;
    if (!ok) {
        std::cerr << "Warning: could not fully overwrite " << path << std::endl;
    }
    if (unlink(path) != 0) {
        std::cerr << "Error deleting " << path << ": " << std::strerror(errno) << std::endl;
        return false;
    }
    return ok;
}

static bool deriveKey(const std::string& passphrase, const unsigned char* salt,
                      unsigned char* key) {
    if (PKCS5_PBKDF2_HMAC(passphrase.data(), static_cast<int>(passphrase.size()),
                          salt, static_cast<int>(SALT_LEN), PBKDF2_ITERATIONS,
                          EVP_sha256(), static_cast<int>(KEY_LEN), key) != 1) {
        printOpenSSLErrors();
        return false;
    }
    return true;
}

static bool encryptFile(const std::string& passphrase) {
    if (access(ENC_FILE, F_OK) == 0) {
        std::cerr << ENC_FILE << " already exists; decrypt or move it first so it is not overwritten." << std::endl;
        return false;
    }

    std::vector<unsigned char> plain;
    if (!readFile(PLAIN_FILE, plain)) return false;
    if (plain.size() > static_cast<size_t>(INT32_MAX) - 64) {
        std::cerr << "File too large" << std::endl;
        wipe(plain);
        return false;
    }

    std::vector<unsigned char> out(HEADER_LEN + plain.size() + TAG_LEN);
    unsigned char* p = out.data();
    std::memcpy(p, MAGIC, sizeof(MAGIC));
    p[sizeof(MAGIC)] = FORMAT_VERSION;
    unsigned char* salt = p + sizeof(MAGIC) + 1;
    unsigned char* nonce = salt + SALT_LEN;
    unsigned char* ct = p + HEADER_LEN;

    unsigned char key[KEY_LEN];
    bool ok = false;
    EVP_CIPHER_CTX* ctx = nullptr;
    int len = 0, total = 0;

    if (RAND_bytes(salt, SALT_LEN) != 1 || RAND_bytes(nonce, NONCE_LEN) != 1) {
        printOpenSSLErrors();
        goto done;
    }
    if (!deriveKey(passphrase, salt, key)) goto done;

    ctx = EVP_CIPHER_CTX_new();
    if (!ctx) { printOpenSSLErrors(); goto done; }
    if (EVP_EncryptInit_ex(ctx, EVP_aes_256_gcm(), nullptr, nullptr, nullptr) != 1 ||
        EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_SET_IVLEN, static_cast<int>(NONCE_LEN), nullptr) != 1 ||
        EVP_EncryptInit_ex(ctx, nullptr, nullptr, key, nonce) != 1 ||
        EVP_EncryptUpdate(ctx, nullptr, &len, out.data(), static_cast<int>(HEADER_LEN)) != 1) {
        printOpenSSLErrors();
        goto done;
    }
    if (!plain.empty()) {
        if (EVP_EncryptUpdate(ctx, ct, &len, plain.data(), static_cast<int>(plain.size())) != 1) {
            printOpenSSLErrors();
            goto done;
        }
        total = len;
    }
    if (EVP_EncryptFinal_ex(ctx, ct + total, &len) != 1) {
        printOpenSSLErrors();
        goto done;
    }
    total += len;
    if (static_cast<size_t>(total) != plain.size() ||
        EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_GET_TAG, static_cast<int>(TAG_LEN), ct + total) != 1) {
        printOpenSSLErrors();
        goto done;
    }

    if (!writeFileAtomic(ENC_FILE, out)) goto done;
    ok = true;

done:
    EVP_CIPHER_CTX_free(ctx);
    OPENSSL_cleanse(key, sizeof(key));
    wipe(plain);
    wipe(out);
    if (!ok) return false;

    // Only now, with the encrypted file safely on disk, remove the plaintext.
    if (!shredFile(PLAIN_FILE)) {
        std::cerr << "Encrypted file written, but the plaintext file was not fully removed. Remove it by hand." << std::endl;
        return false;
    }
    return true;
}

static bool decryptFile(const std::string& passphrase) {
    if (access(PLAIN_FILE, F_OK) == 0) {
        std::cerr << PLAIN_FILE << " already exists; move it first so it is not overwritten." << std::endl;
        return false;
    }

    std::vector<unsigned char> in;
    if (!readFile(ENC_FILE, in)) return false;
    if (in.size() < HEADER_LEN + TAG_LEN ||
        std::memcmp(in.data(), MAGIC, sizeof(MAGIC)) != 0) {
        std::cerr << ENC_FILE << " is not an ObsidianDart file (or was made by the old, unsupported format)." << std::endl;
        return false;
    }
    if (in[sizeof(MAGIC)] != FORMAT_VERSION) {
        std::cerr << "Unsupported file format version " << static_cast<int>(in[sizeof(MAGIC)]) << std::endl;
        return false;
    }
    if (in.size() - HEADER_LEN - TAG_LEN > static_cast<size_t>(INT32_MAX) - 64) {
        std::cerr << "File too large" << std::endl;
        return false;
    }

    const unsigned char* salt = in.data() + sizeof(MAGIC) + 1;
    const unsigned char* nonce = salt + SALT_LEN;
    const unsigned char* ct = in.data() + HEADER_LEN;
    const size_t ctLen = in.size() - HEADER_LEN - TAG_LEN;
    unsigned char tag[TAG_LEN];
    std::memcpy(tag, ct + ctLen, TAG_LEN);

    std::vector<unsigned char> plain(ctLen);
    unsigned char key[KEY_LEN];
    bool ok = false;
    EVP_CIPHER_CTX* ctx = nullptr;
    int len = 0, total = 0;

    if (!deriveKey(passphrase, salt, key)) goto done;

    ctx = EVP_CIPHER_CTX_new();
    if (!ctx) { printOpenSSLErrors(); goto done; }
    if (EVP_DecryptInit_ex(ctx, EVP_aes_256_gcm(), nullptr, nullptr, nullptr) != 1 ||
        EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_SET_IVLEN, static_cast<int>(NONCE_LEN), nullptr) != 1 ||
        EVP_DecryptInit_ex(ctx, nullptr, nullptr, key, nonce) != 1 ||
        EVP_DecryptUpdate(ctx, nullptr, &len, in.data(), static_cast<int>(HEADER_LEN)) != 1) {
        printOpenSSLErrors();
        goto done;
    }
    if (ctLen > 0) {
        if (EVP_DecryptUpdate(ctx, plain.data(), &len, ct, static_cast<int>(ctLen)) != 1) {
            printOpenSSLErrors();
            goto done;
        }
        total = len;
    }
    if (EVP_CIPHER_CTX_ctrl(ctx, EVP_CTRL_GCM_SET_TAG, static_cast<int>(TAG_LEN), tag) != 1) {
        printOpenSSLErrors();
        goto done;
    }
    // The tag is checked here. Nothing is written unless it matches.
    if (EVP_DecryptFinal_ex(ctx, plain.data() + total, &len) != 1) {
        std::cerr << "Decryption failed: wrong passphrase, or the file has been modified or corrupted." << std::endl;
        goto done;
    }
    total += len;
    if (static_cast<size_t>(total) != ctLen) goto done;

    if (!writeFileAtomic(PLAIN_FILE, plain)) goto done;
    ok = true;

done:
    EVP_CIPHER_CTX_free(ctx);
    OPENSSL_cleanse(key, sizeof(key));
    wipe(plain);
    wipe(in);
    return ok;
}

static const char* const VERSION = "o1.2A";

static void printUsage() {
    std::cout << "ObsidianDart " << VERSION << ": encrypt or decrypt " << PLAIN_FILE
              << " in the current folder.\n\n"
              << "Usage: obsidiandart [-h | --help] [-V | --version]\n\n"
              << "Run it in the folder that holds " << PLAIN_FILE << " (to encrypt)\n"
              << "or " << ENC_FILE << " (to decrypt) and follow the menu.\n"
              << "Set NO_COLOR to turn colors off." << std::endl;
}

static bool exists(const char* path) { return access(path, F_OK) == 0; }

static void printStatus(bool havePlain, bool haveEnc) {
    char cwd[4096];
    const char* where = getcwd(cwd, sizeof(cwd)) ? cwd : ".";
    std::cout << "\n  " << BOLD() << CYAN() << "ObsidianDart" << RESET() << " "
              << DIM() << VERSION << "  AES-256-GCM file vault" << RESET() << "\n"
              << "  " << DIM() << "------------------------------------------" << RESET() << "\n"
              << "  folder  " << where << "\n";
    auto line = [](const char* name, bool here, const char* label) {
        std::cout << "  " << (here ? GREEN() : DIM()) << (here ? "  [x] " : "  [ ] ")
                  << name << std::string(std::strlen(ENC_FILE) - std::strlen(name), ' ')
                  << RESET() << DIM() << "  " << label << RESET() << "\n";
    };
    line(PLAIN_FILE, havePlain, havePlain ? "plaintext, unprotected" : "not here");
    line(ENC_FILE, haveEnc, haveEnc ? "encrypted" : "not here");
    std::cout << std::endl;
}

// Shows the menu until a valid choice or EOF. Returns '1', '2' or 'q'.
// Enter alone picks the suggested action for the files that are present.
static char askChoice(bool havePlain, bool haveEnc) {
    char suggested = 0;
    if (havePlain && !haveEnc) suggested = '1';
    else if (haveEnc && !havePlain) suggested = '2';

    std::cout << "  " << BOLD() << "1" << RESET() << "  Encrypt " << PLAIN_FILE << "\n"
              << "  " << BOLD() << "2" << RESET() << "  Decrypt " << ENC_FILE << "\n"
              << "  " << BOLD() << "q" << RESET() << "  Quit\n" << std::endl;
    for (;;) {
        std::cout << "  Choice" << (suggested ? std::string(" [") + suggested + "]" : std::string())
                  << ": " << std::flush;
        std::string choice;
        if (!std::getline(std::cin, choice)) { std::cout << std::endl; return 'q'; }
        if (choice.empty() && suggested) return suggested;
        if (choice == "1" || choice == "2") return choice[0];
        if (choice == "q" || choice == "Q") return 'q';
        std::cout << "  " << YELLOW() << "Type 1, 2 or q." << RESET() << std::endl;
    }
}

int main(int argc, char** argv) {
    for (int i = 1; i < argc; ++i) {
        std::string a = argv[i];
        if (a == "-h" || a == "--help") { printUsage(); return 0; }
        if (a == "-V" || a == "--version") { std::cout << "obsidiandart " << VERSION << std::endl; return 0; }
        fail("unknown option " + a + " (try --help)");
        return 2;
    }

    signal(SIGINT, restoreTermOnSignal);
    signal(SIGTERM, restoreTermOnSignal);

    bool havePlain = exists(PLAIN_FILE);
    bool haveEnc = exists(ENC_FILE);
    printStatus(havePlain, haveEnc);

    char choice = askChoice(havePlain, haveEnc);
    if (choice == 'q') return 0;

    // Check the files before asking for a passphrase, so nobody types it in
    // for nothing.
    if (choice == '1' && !havePlain) {
        fail(std::string("no ") + PLAIN_FILE + " in this folder to encrypt.");
        return 1;
    }
    if (choice == '2' && !haveEnc) {
        fail(std::string("no ") + ENC_FILE + " in this folder to decrypt.");
        return 1;
    }

    std::cout << std::endl;
    std::string passphrase;
    if (!readSecret("  Passphrase: ", passphrase) || passphrase.empty()) {
        fail("no passphrase entered.");
        return 1;
    }

    bool ok = false;
    if (choice == '1') {
        std::string confirm;
        if (!readSecret("  Confirm passphrase: ", confirm) || confirm != passphrase) {
            fail("passphrases do not match. Nothing was changed.");
            wipe(confirm);
            wipe(passphrase);
            return 1;
        }
        wipe(confirm);
        std::cout << "  " << DIM() << "Deriving key and encrypting..." << RESET() << std::endl;
        ok = encryptFile(passphrase);
        if (ok) {
            std::cout << "  " << GREEN() << "Encrypted." << RESET() << " " << ENC_FILE
                      << " written, " << PLAIN_FILE << " overwritten and removed." << std::endl;
        }
    } else {
        std::cout << "  " << DIM() << "Deriving key and decrypting..." << RESET() << std::endl;
        ok = decryptFile(passphrase);
        // Same as before: the encrypted copy is removed after a successful
        // decrypt. It is kept if anything went wrong.
        if (ok) {
            std::cout << "  " << GREEN() << "Decrypted." << RESET() << " " << PLAIN_FILE
                      << " written." << std::endl;
            if (unlink(ENC_FILE) != 0) {
                fail(std::string("could not delete ") + ENC_FILE + ": " + std::strerror(errno));
            }
            std::cout << "  " << YELLOW() << "Note:" << RESET() << " " << PLAIN_FILE
                      << " is plaintext on disk until you encrypt it again." << std::endl;
        }
    }

    wipe(passphrase);
    if (!ok) std::cerr << "  " << RED() << "Failed." << RESET() << " Nothing else was changed." << std::endl;
    return ok ? 0 : 1;
}
