// 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 <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);
}

// Reads one line from stdin with terminal echo turned off (when stdin is a
// terminal). Echo is always restored before returning.
static bool readSecret(const std::string& prompt, std::string& out) {
    std::cout << prompt << std::flush;
    struct termios oldFlags;
    bool isTty = (tcgetattr(STDIN_FILENO, &oldFlags) == 0);
    if (isTty) {
        struct termios newFlags = oldFlags;
        newFlags.c_lflag &= ~ECHO;
        tcsetattr(STDIN_FILENO, TCSANOW, &newFlags);
    }
    bool ok = static_cast<bool>(std::getline(std::cin, out));
    if (isTty) {
        tcsetattr(STDIN_FILENO, TCSANOW, &oldFlags);
    }
    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;
}

int main() {
    std::string passphrase;
    if (!readSecret("Enter passphrase: ", passphrase) || passphrase.empty()) {
        std::cerr << "No passphrase entered" << std::endl;
        return 1;
    }

    std::cout << "1. Encrypt file" << std::endl;
    std::cout << "2. Decrypt file" << std::endl;
    std::cout << "Enter your choice: " << std::flush;
    std::string choice;
    std::getline(std::cin, choice);

    bool ok = false;
    if (choice == "1") {
        std::string confirm;
        if (!readSecret("Confirm passphrase: ", confirm) || confirm != passphrase) {
            std::cerr << "Passphrases do not match" << std::endl;
            wipe(confirm);
            wipe(passphrase);
            return 1;
        }
        wipe(confirm);
        std::cout << "Encrypting file..." << std::endl;
        ok = encryptFile(passphrase);
    } else if (choice == "2") {
        std::cout << "Decrypting file..." << 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 && unlink(ENC_FILE) != 0) {
            std::cerr << "Error deleting encrypted file" << std::endl;
        }
    } else {
        std::cerr << "Invalid choice" << std::endl;
    }

    wipe(passphrase);
    if (ok) std::cout << "Done." << std::endl;
    return ok ? 0 : 1;
}
