#!/usr/bin/env python3
"""
syncture-exit.py — turn a copy of a Syncture bucket back into Revit models.

Runs against a folder that holds a copy of the bucket your workspace was kept
in, and writes every version of every model as the files that were published.
Needs nothing from Syncture: no account, no service, no network.

    python syncture-exit.py --bucket .\\bucket --out .\\models
    python syncture-exit.py --bucket .\\bucket --out .\\models --key SYNK1-...
    python syncture-exit.py --bucket .\\bucket --out .\\models --latest
    python syncture-exit.py --bucket .\\bucket --verify

--bucket   A folder holding a copy of the bucket, or of the workspace's own
           prefix (the folder that contains `b/` and `v/`). Make the copy with
           the tool you already use: aws s3 sync, rclone sync, mc mirror.
--out      Where the models are written: <out>/<workspace>/<model path>/v0002/...
--key      The workspace key, for a workspace created with encryption on, or
--key-file a file holding it. Never needed for any other workspace.
--latest   Only the newest version of each model.
--model    One model, by its id or its name.
--verify   Check every object, every difference and every file, and write nothing.

Python 3.8 or newer. The standard library does everything except decrypt: a
workspace created with encryption on needs `pip install cryptography`.

The layout and both byte formats are documented at
https://syncture.com/docs/exit and in md/EXIT.md of the Syncture repository.
Exit status is 0 when every version was rebuilt and checked, 1 otherwise.
"""

import argparse
import hashlib
import hmac
import json
import os
import struct
import sys
import zlib

FORMAT = "syncture-version/1"
ALPHABET = "0123456789ABCDEFGHJKMNPQRSTVWXYZ"


class ExitError(Exception):
    """A reason one version could not be rebuilt. Named, and never silent."""


# ---------------------------------------------------------------------------
# The key: SYNK1- and seven groups of Crockford base32
# ---------------------------------------------------------------------------

def parse_key(text):
    """The 32 key bytes of a SYNK1 key string, or an ExitError naming what is wrong."""
    cleaned = "".join(ch for ch in text.replace("﻿", "") if not ch.isspace()).upper()
    if not cleaned.startswith("SYNK1-"):
        raise ExitError("not a Syncture key: it does not start with SYNK1-")
    folded = cleaned[6:].replace("-", "").replace("I", "1").replace("L", "1").replace("O", "0")
    for ch in folded:
        if ch not in ALPHABET:
            raise ExitError("not a Syncture key: the character %r is not in its alphabet" % ch)
    if len(folded) != 56:
        raise ExitError("not a Syncture key: %d characters after the prefix, not 56" % len(folded))

    key = bytearray()
    accumulator = 0
    bits = 0
    for ch in folded[:52]:
        accumulator = ((accumulator << 5) | ALPHABET.index(ch)) & 0xFFFF
        bits += 5
        if bits >= 8:
            bits -= 8
            if len(key) < 32:
                key.append((accumulator >> bits) & 0xFF)
    if len(key) != 32 or accumulator & ((1 << bits) - 1):
        raise ExitError("not a Syncture key: its padding is not zero")

    digest = hashlib.sha256(b"syncture/key/v1" + bytes(key)).digest()
    check = "".join(format(b, "08b") for b in digest[:3])
    checksum = "".join(ALPHABET[int(check[i:i + 5], 2)] for i in range(0, 20, 5))
    if checksum != folded[52:]:
        raise ExitError("not a Syncture key: it does not check")
    return bytes(key)


def key_id_of(key):
    return hashlib.sha256(b"syncture/kid/v1" + key).digest()[:8].hex()


def subkeys(key, workspace_id):
    """RFC 5869 HKDF-SHA256, written out: one extract, one expand per use."""
    prk = hmac.new(b"syncture/workspace/v1", key, hashlib.sha256).digest()

    def expand(use):
        return hmac.new(prk, ("%s|%s" % (use, workspace_id)).encode("utf-8") + b"\x01", hashlib.sha256).digest()

    return {"enc": expand("enc"), "mac": expand("mac"), "file": expand("file")}


# ---------------------------------------------------------------------------
# The envelope: SYNK, version 1
# ---------------------------------------------------------------------------

def open_envelope(keys, key_id, envelope):
    """The plaintext of a sealed object: checked, in order, before one block is decrypted."""
    if len(envelope) < 77:
        raise ExitError("a sealed object is too short")
    if envelope[:4] != b"SYNK" or envelope[4] != 1:
        raise ExitError("an object is not in Syncture's sealed format")
    if (len(envelope) - 61) % 16:
        raise ExitError("a sealed object has an impossible length")
    named = envelope[5:13].hex()
    if named != key_id:
        raise ExitError("an object was sealed under key %s, and this key is %s" % (named, key_id))
    body, tag = envelope[:-32], envelope[-32:]
    if not hmac.compare_digest(hmac.new(keys["mac"], body, hashlib.sha256).digest(), tag):
        raise ExitError("a sealed object failed its authentication check")

    try:
        from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
    except ImportError:
        raise ExitError("this workspace was created with encryption on, and decrypting needs "
                        "`pip install cryptography`")
    decryptor = Cipher(algorithms.AES(keys["enc"]), modes.CBC(envelope[13:29])).decryptor()
    padded = decryptor.update(envelope[29:-32]) + decryptor.finalize()
    pad = padded[-1] if padded else 0
    if pad < 1 or pad > 16 or padded[-pad:] != bytes([pad]) * pad:
        raise ExitError("a sealed object has bad padding")
    return padded[:-pad]


# ---------------------------------------------------------------------------
# The difference: SYND, version 1
# ---------------------------------------------------------------------------

def decode_delta(delta, baseline):
    """The chunk a difference stands for, rebuilt against the bases' bytes."""
    if len(delta) < 14 or delta[:4] != b"SYND" or delta[4] != 1:
        raise ExitError("an object named as a difference is not one")
    plain_length, base_length = struct.unpack_from("<II", delta, 5)
    base_count = delta[13]
    if base_length != len(baseline):
        raise ExitError("a difference was taken against other bases than the ones named")
    if not 1 <= base_count <= 3:
        raise ExitError("a difference names %d bases" % base_count)

    at = 14 + 32 * base_count
    if len(delta) < at + 4:
        raise ExitError("a difference is cut short")
    (operations,) = struct.unpack_from("<I", delta, at)
    at += 4
    ops = []
    literal_length = 0
    produced = 0
    for _ in range(operations):
        if len(delta) < at + 9:
            raise ExitError("a difference is cut short")
        kind, a, b = struct.unpack_from("<BII", delta, at)
        at += 9
        if kind == 0:
            if b == 0 or a + b > len(baseline):
                raise ExitError("a difference copies from outside its bases")
        elif kind == 1:
            if a != 0:
                raise ExitError("a difference is malformed")
            literal_length += b
        else:
            raise ExitError("a difference holds an unknown operation")
        produced += b
        ops.append((kind, a, b))
    if produced != plain_length:
        raise ExitError("a difference's operations do not produce its stated length")

    if len(delta) < at + 4:
        raise ExitError("a difference is cut short")
    (packed_length,) = struct.unpack_from("<I", delta, at)
    at += 4
    if len(delta) - at != packed_length:
        raise ExitError("a difference's literal block is not the stated length")
    try:
        literals = zlib.decompress(delta[at:], -15)
    except zlib.error:
        raise ExitError("a difference's literal block does not inflate")
    if len(literals) != literal_length:
        raise ExitError("a difference's literal block inflates to the wrong length")

    out = bytearray(plain_length)
    into = 0
    taken = 0
    for kind, a, b in ops:
        if kind == 0:
            out[into:into + b] = baseline[a:a + b]
        else:
            out[into:into + b] = literals[taken:taken + b]
            taken += b
        into += b
    return bytes(out)


# ---------------------------------------------------------------------------
# The bucket
# ---------------------------------------------------------------------------

def workspace_roots(bucket):
    """Every folder under `bucket` that holds a workspace: one with `v/` and `b/` inside."""
    bucket = os.path.abspath(bucket)
    if os.path.isdir(os.path.join(bucket, "v")):
        return [bucket]
    roots = []
    for dirpath, dirnames, _ in os.walk(bucket):
        if os.path.relpath(dirpath, bucket).count(os.sep) > 4:
            dirnames[:] = []
            continue
        if "v" in dirnames and os.path.basename(os.path.dirname(dirpath)) == "w":
            roots.append(dirpath)
            dirnames[:] = []
    return sorted(roots)


def read_objects(root):
    """Every version object under `v/`, parsed and checked for its format."""
    versions = []
    base = os.path.join(root, "v")
    for model_dir in sorted(os.listdir(base)):
        folder = os.path.join(base, model_dir)
        if not os.path.isdir(folder):
            continue
        for name in sorted(os.listdir(folder)):
            if not name.endswith(".json"):
                continue
            with open(os.path.join(folder, name), "rb") as handle:
                try:
                    obj = json.loads(handle.read().decode("utf-8"))
                except (ValueError, UnicodeDecodeError):
                    raise ExitError("%s is not a version object" % os.path.join(folder, name))
            if obj.get("format") != FORMAT:
                raise ExitError("%s is in format %r, and this script reads %s" % (name, obj.get("format"), FORMAT))
            versions.append(obj)
    return versions


def safe_name(segment):
    """A folder or file name Windows and POSIX both accept."""
    cleaned = "".join("_" if ch in '<>:"/\\|?*' or ord(ch) < 32 else ch for ch in segment).rstrip(". ")
    return cleaned or "_"


class Rebuilder:
    def __init__(self, root, key_text, out, verify_only):
        self.root = root
        self.out = out
        self.verify_only = verify_only
        self.key = parse_key(key_text) if key_text else None
        self.key_id = key_id_of(self.key) if self.key else ""
        self.subkeys = {}
        self.checked = {}

    def stored(self, digest):
        """A stored object by its digest, re-hashed against its name before it is trusted."""
        path = os.path.join(self.root, "b", digest)
        if not os.path.isfile(path):
            raise ExitError("object b/%s is not in the bucket" % digest)
        with open(path, "rb") as handle:
            data = handle.read()
        if digest not in self.checked:
            if hashlib.sha256(data).hexdigest() != digest:
                raise ExitError("object b/%s does not hash to its name; the copy is damaged" % digest)
            self.checked[digest] = True
        return data

    def plaintext(self, keys, digest):
        data = self.stored(digest)
        return open_envelope(keys, self.key_id, data) if keys else data

    def keys_for(self, obj):
        if not obj.get("encryption"):
            return None
        if not self.key:
            raise ExitError("this workspace was created with encryption on; pass its key with --key")
        wanted = obj["encryption"]["keyId"]
        if wanted != self.key_id:
            raise ExitError("this workspace's key has id %s, and the key given has id %s" % (wanted, self.key_id))
        workspace_id = obj["workspace"]["id"]
        if workspace_id not in self.subkeys:
            self.subkeys[workspace_id] = subkeys(self.key, workspace_id)
        return self.subkeys[workspace_id]

    def rebuild_file(self, keys, entry):
        """One file's bytes, checked at every step."""
        chunks = entry.get("chunks", "")
        if len(chunks) % 64:
            raise ExitError("%s: the chunk list is not a list of digests" % entry["path"])
        digests = [chunks[i:i + 64] for i in range(0, len(chunks), 64)]
        lens = [int(x) for x in entry["lens"].split(",")] if entry.get("lens") else []
        if len(lens) != len(digests):
            raise ExitError("%s: %d lengths for %d chunks" % (entry["path"], len(lens), len(digests)))
        if sum(lens) != entry["size"]:
            raise ExitError("%s: the chunk lengths do not add up to the file's size" % entry["path"])
        deltas = {d["i"]: d for d in entry.get("deltas") or []}

        parts = []
        for index, digest in enumerate(digests):
            delta = deltas.get(index)
            if delta:
                baseline = b"".join(self.plaintext(keys, base) for base in delta["bases"])
                content = decode_delta(self.plaintext(keys, digest), baseline)
                if hashlib.sha256(content).hexdigest() != delta["full"]:
                    raise ExitError("%s: chunk %d rebuilt to the wrong bytes" % (entry["path"], index))
            else:
                content = self.plaintext(keys, digest)
            if len(content) != lens[index]:
                raise ExitError("%s: chunk %d is %d bytes, not %d" % (entry["path"], index, len(content), lens[index]))
            parts.append(content)

        data = b"".join(parts)
        if entry.get("sha256"):
            if keys:
                digest = hmac.new(keys["file"], data, hashlib.sha256).hexdigest()
            else:
                digest = hashlib.sha256(data).hexdigest()
            if digest != entry["sha256"]:
                raise ExitError("%s: the whole file does not check" % entry["path"])
        return data

    def write(self, folder, entry, data):
        target = os.path.join(folder, *[safe_name(part) for part in entry["path"].split("/")])
        os.makedirs(os.path.dirname(target), exist_ok=True)
        partial = target + ".partial"
        with open(partial, "wb") as handle:
            handle.write(data)
        os.replace(partial, target)


def main(argv):
    parser = argparse.ArgumentParser(description="Turn a copy of a Syncture bucket back into Revit models.")
    parser.add_argument("--bucket", required=True, help="a folder holding a copy of the bucket, or of the workspace's prefix")
    parser.add_argument("--out", help="where the models are written")
    parser.add_argument("--key", help="the workspace key, for a workspace created with encryption on")
    parser.add_argument("--key-file", help="a file holding the workspace key")
    parser.add_argument("--latest", action="store_true", help="only the newest version of each model")
    parser.add_argument("--model", help="one model, by id or name")
    parser.add_argument("--verify", action="store_true", help="check everything and write nothing")
    args = parser.parse_args(argv)

    if not args.verify and not args.out:
        parser.error("--out is needed unless --verify is given")
    key_text = args.key
    if args.key_file:
        with open(args.key_file, "r", encoding="utf-8-sig") as handle:
            key_text = handle.read()

    roots = workspace_roots(args.bucket)
    if not roots:
        print("No workspace found under %s: no folder holding v/ and b/." % args.bucket)
        return 1

    failures = 0
    rebuilt = 0
    files = 0
    total = 0
    for root in roots:
        try:
            worker = Rebuilder(root, key_text, args.out, args.verify)
            versions = read_objects(root)
        except ExitError as error:
            print("x  %s: %s" % (root, error))
            failures += 1
            continue
        if not versions:
            print("-  %s holds no versions" % root)
            continue

        by_model = {}
        for obj in versions:
            by_model.setdefault(obj["model"]["id"], []).append(obj)
        for model_id, objs in sorted(by_model.items()):
            objs.sort(key=lambda o: o["version"]["seq"])
            newest = objs[-1]
            if args.model and args.model not in (model_id, newest["model"]["name"]):
                continue
            chosen = objs[-1:] if args.latest else objs
            workspace_name = safe_name(newest["workspace"]["name"] or newest["workspace"]["id"])
            model_folder = [safe_name(part) for part in newest["model"]["path"].split("/") if part]
            print("%s  (%s), %d of %d versions" % (newest["model"]["path"], model_id, len(chosen), len(objs)))

            for obj in chosen:
                seq = obj["version"]["seq"]
                try:
                    keys = worker.keys_for(obj)
                    folder = None
                    if not args.verify:
                        folder = os.path.join(args.out, workspace_name, *model_folder, "v%04d" % seq)
                    for entry in obj["files"]:
                        data = worker.rebuild_file(keys, entry)
                        if folder:
                            worker.write(folder, entry, data)
                        files += 1
                        total += len(data)
                    rebuilt += 1
                    print("   +  v%d  %s  %s  %d files" % (seq, obj["version"]["createdAt"][:10], obj["version"]["author"], len(obj["files"])))
                except ExitError as error:
                    failures += 1
                    print("   x  v%d: %s" % (seq, error))

    verb = "checked" if args.verify else "rebuilt"
    print("%d versions %s, %d files, %d bytes, %d failures." % (rebuilt, verb, files, total, failures))
    return 1 if failures else 0


if __name__ == "__main__":
    sys.exit(main(sys.argv[1:]))
