Newer
Older
gnexus-creds / scripts / migrate-crypto.py
"""Crypto migration: scrypt master derivation + AAD-bound field ciphertexts.

Re-encrypts two layers in the live database:
  1. every wrapped user key (user_encryption_keys) — unwrapped with the
     legacy sha256 master and re-wrapped with the salted scrypt derivation;
  2. every encrypted field (secret_versions.fields) — decrypted and
     re-encrypted with an aad binding it to "<user_id>:<secret_id>:<version_id>".

Idempotent: already-migrated items are detected and skipped, so re-running
after an interruption only finishes the rest. Requires the NEW salt to be
set (GNEXUS_CREDS_MASTER_KEY_SALT); the master key itself must be the same
env value as before the migration.

Run from the project root with the project venv:
    .venv/bin/python scripts/migrate-crypto.py [--dry-run]
"""

from __future__ import annotations

import argparse
import sys

from cryptography.exceptions import InvalidTag

from gnexus_creds import crypto
from gnexus_creds.crypto import _b64d
from sqlalchemy.orm.attributes import flag_modified
from gnexus_creds.config import get_settings
from gnexus_creds.db import SessionLocal
from gnexus_creds.models import Secret, UserEncryptionKey


def _rebind_version_fields(key: bytes, version, aad: str, dry: bool) -> int:
    """Return how many ciphertexts were rebound (0 for error rows)."""
    rebound = 0
    for field in version.fields:
        if not field.get("encrypted"):
            continue
        envelope = field["value"]
        # not crypto.decrypt_text here: its legacy fallback would accept the
        # unbound form as "bound" and skip the rebind — probe auth directly
        nonce, ciphertext = _b64d(envelope["nonce"]), _b64d(envelope["ciphertext"])
        try:
            crypto.decrypt_bytes(key, nonce, ciphertext, aad=aad.encode())
            continue  # already bound
        except InvalidTag:
            pass
        try:
            plaintext = crypto.decrypt_bytes(key, nonce, ciphertext).decode()
        except InvalidTag:
            print("   !! cannot decrypt an encrypted field — left as is", file=sys.stderr)
            continue
        if not dry:
            field["value"] = crypto.encrypt_text(key, envelope["key_id"], plaintext, aad=aad)
        rebound += 1
    return rebound


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--dry-run", action="store_true", help="report without writing")
    args = parser.parse_args()

    settings = get_settings()
    if not settings.master_key_salt:
        print(
            "GNEXUS_CREDS_MASTER_KEY_SALT is not set — nothing to migrate to.",
            file=sys.stderr,
        )
        return 2
    new_master = crypto.derive_master_key(settings.master_key, settings.master_key_salt)
    legacy_master = crypto._derive_legacy(settings.master_key)

    db = SessionLocal()
    keys_rewrapped = keys_ok = fields_rebound = 0

    # 1. user keys: legacy-sha256-wrapped -> scrypt-wrapped
    rows = list(db.query(UserEncryptionKey))
    keys = {}
    for row in rows:
        try:
            # already wrapped under the new derivation
            keys[row.user_id] = crypto.decrypt_bytes(new_master, row.nonce, row.encrypted_key)
            keys_ok += 1
            continue
        except InvalidTag:
            pass
        try:
            keys[row.user_id] = crypto.decrypt_bytes(legacy_master, row.nonce, row.encrypted_key)
        except InvalidTag:
            print(
                f"!! cannot unwrap key {row.key_id} (user {row.user_id}) with either master"
                " — skipped, its data stays untouched",
                file=sys.stderr,
            )
            continue
        if not args.dry_run:
            nonce, encrypted_key = crypto.encrypt_bytes(new_master, keys[row.user_id])
            row.nonce = nonce
            row.encrypted_key = encrypted_key
        keys_rewrapped += 1
    if not args.dry_run:
        db.commit()

    # 2. field ciphertexts: unbound legacy -> aad "<user:secret:version>"
    for secret in db.query(Secret):
        key = keys.get(secret.user_id)
        if key is None:
            continue
        for version in secret.versions:
            aad = f"{secret.user_id}:{secret.id}:{version.id}"
            rebound = _rebind_version_fields(key, version, aad, args.dry_run)
            if rebound and not args.dry_run:
                # equal-content reassignment is not a change for SQLAlchemy's
                # history — force the JSON blob flush properly
                flag_modified(version, "fields")
            fields_rebound += rebound
    if not args.dry_run:
        db.commit()

    mode = "DRY RUN: " if args.dry_run else ""
    print(
        f"{mode}keys rewrapped {keys_rewrapped} (already ok: {keys_ok}); "
        f"field ciphertexts rebound {fields_rebound}"
    )
    return 0


if __name__ == "__main__":
    raise SystemExit(main())