#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
SDX 命令行工具 —— 文件加密 BitCipher · 星辉数盾 的跨端实现。
与浏览器扩展读写同一种 .sdx 容器：写 v2（Argon2id + 分块 AES-256-GCM），读 v1 / v2。

    python3 sdx.py enc  文件 [-o 输出.sdx] [-p 口令]
    python3 sdx.py dec  文件.sdx [-o 输出] [-p 口令]
    python3 sdx.py enc-text "明文"        → 输出 SDX:base64
    python3 sdx.py dec-text "SDX:..."     → 输出明文
    python3 sdx.py info 文件.sdx          → 打印容器头（不需要口令）

口令也可以用环境变量 SDX_PASSWORD，或不传 -p 交互输入。
依赖：pip install cryptography argon2-cffi
"""
import argparse
import base64
import getpass
import json
import os
import struct
import sys
import unicodedata

try:
    from cryptography.hazmat.primitives.ciphers.aead import AESGCM
    from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
    from cryptography.hazmat.primitives import hashes
    from argon2.low_level import hash_secret_raw, Type
except ImportError:
    print("缺少依赖，请运行: pip install cryptography argon2-cffi", file=sys.stderr)
    sys.exit(1)

MAGIC = b"SDX1"
VER_1, VER_2 = 1, 2
KDF = {"name": "argon2id", "m": 65536, "t": 3, "p": 1}
CHUNK = 1 << 20
TAG = 16
TEXT_PREFIX = "SDX:"
MAX_HEADER = 4096
LIMITS = dict(argon2_m=1 << 20, argon2_t=20, argon2_p=16, pbkdf2_iters=10_000_000, chunk_min=1024, chunk_max=16 << 20, meta_max=4096)


class SDXError(Exception):
    pass


# ---------- 头 ----------

def pack_header(header, ver=VER_2):
    hjson = json.dumps(header, separators=(",", ":"), ensure_ascii=False).encode("utf-8")
    if len(hjson) > MAX_HEADER:
        raise SDXError("SDX header too large")
    return MAGIC + bytes([ver]) + struct.pack("<H", len(hjson)) + hjson


def parse_header(buf):
    """返回 (ver, header, prefix_len)；字节不够返回 None。"""
    if len(buf) < 7:
        return None
    if buf[:4] != MAGIC:
        raise SDXError("不是 SDX 文件（魔数不匹配）")
    ver = buf[4]
    if ver not in (VER_1, VER_2):
        raise SDXError("不支持的 SDX 版本 %d" % ver)
    hlen = struct.unpack("<H", buf[5:7])[0]
    if hlen > MAX_HEADER:
        raise SDXError("SDX header too large")
    if len(buf) < 7 + hlen:
        return None
    try:
        header = json.loads(buf[7:7 + hlen].decode("utf-8"))
    except Exception:
        raise SDXError("SDX 头部损坏")
    if not isinstance(header, dict):
        raise SDXError("SDX 头部损坏")
    return ver, header, 7 + hlen


# ---------- 密钥 ----------

def derive_key(password, header):
    salt = base64.b64decode(header.get("salt", ""))
    if len(salt) < 8:
        raise SDXError("SDX bad salt")
    k = header.get("kdf") or {}
    name = k.get("name")
    if name == "argon2id":
        m, t, p = int(k.get("m", 0)), int(k.get("t", 0)), int(k.get("p", 0))
        if not (8 * p <= m <= LIMITS["argon2_m"] and 1 <= t <= LIMITS["argon2_t"] and 1 <= p <= LIMITS["argon2_p"]):
            raise SDXError("SDX kdf params out of range")
        # v2 口令 NFC 归一，与扩展一致（v1 保持原样）
        pw = unicodedata.normalize("NFC", password).encode("utf-8")
        return hash_secret_raw(pw, salt, time_cost=t, memory_cost=m, parallelism=p, hash_len=32, type=Type.ID)
    if name == "PBKDF2-HS256":
        iters = int(k.get("iters", 0))
        if not (1000 <= iters <= LIMITS["pbkdf2_iters"]):
            raise SDXError("SDX kdf params out of range")
        return PBKDF2HMAC(algorithm=hashes.SHA256(), length=32, salt=salt, iterations=iters).derive(password.encode("utf-8"))
    raise SDXError("不支持的 KDF")


def nonce(counter, last):
    return b"\x00\x00\x00" + struct.pack(">Q", counter) + (b"\x01" if last else b"\x00")


# ---------- 分块 ----------

def chunkify(reader, size):
    """reader: 可迭代的 bytes 片段。攒够 > size 才吐块，结束时剩余的当末块。产出 (block, last)。"""
    buf = bytearray()
    for piece in reader:
        if not piece:
            continue
        buf += piece
        while len(buf) > size:
            yield bytes(buf[:size]), False
            del buf[:size]
    yield bytes(buf), True


def read_pieces(fp, piece=1 << 16):
    while True:
        b = fp.read(piece)
        if not b:
            return
        yield b


# ---------- 加密 ----------

def encrypt_stream(reader, password, meta, out, progress=None):
    """reader → out（file-like，二进制写）。返回写出的字节数。"""
    salt = os.urandom(16)
    header = {"alg": "AES-256-GCM", "kdf": dict(KDF), "salt": base64.b64encode(salt).decode(), "chunk": CHUNK, "tag_len": TAG}
    prefix = pack_header(header, VER_2)
    aes = AESGCM(derive_key(password, header))
    out.write(prefix)
    written = len(prefix)
    meta_bytes = json.dumps(meta or {}, separators=(",", ":"), ensure_ascii=False).encode("utf-8")
    if len(meta_bytes) > LIMITS["meta_max"]:
        raise SDXError("SDX meta too large")

    def plain():
        yield struct.pack(">I", len(meta_bytes)) + meta_bytes
        for p in reader:
            yield p

    done = 0
    for counter, (block, last) in enumerate(chunkify(plain(), CHUNK)):
        ct = aes.encrypt(nonce(counter, last), block, prefix)
        out.write(ct)
        written += len(ct)
        done += len(block)
        if progress:
            progress(max(0, done - 4 - len(meta_bytes)))
    return written


# ---------- 解密 ----------

def open_decrypt(reader, password):
    """返回 (ver, header, meta, parts)，parts 是明文片段的生成器。"""
    it = iter(reader)
    head = b""
    parsed = None
    ended = False
    while True:
        parsed = parse_header(head)
        if parsed:
            break
        try:
            head += next(it)
        except StopIteration:
            ended = True
            break
    if not parsed:
        raise SDXError("SDX 文件过短")
    ver, header, plen = parsed

    def rest():
        if len(head) > plen:
            yield head[plen:]
        if ended:
            return
        for p in it:
            yield p

    if ver == VER_1:
        body = head[:plen] + b"".join(rest())
        tag_len = int(header.get("tag_len") or 16)
        if len(body) < plen + tag_len:
            raise SDXError("SDX 文件过短")
        key = derive_key(password, header)
        iv = base64.b64decode(header.get("iv", ""))
        if len(iv) != 12:
            raise SDXError("SDX bad iv")
        try:
            pt = AESGCM(key).decrypt(iv, body[plen:], None)
        except Exception:
            raise SDXError("口令不正确或文件已损坏")
        meta = {"name": header.get("origName", ""), "mime": header.get("mime", ""), "size": len(pt)}
        return ver, header, meta, iter([pt])

    chunk, tag_len = int(header.get("chunk") or 0), int(header.get("tag_len") or 0)
    if header.get("alg") != "AES-256-GCM" or tag_len != TAG:
        raise SDXError("不支持的算法")
    if not (LIMITS["chunk_min"] <= chunk <= LIMITS["chunk_max"]):
        raise SDXError("SDX bad chunk size")
    prefix = head[:plen]
    aes = AESGCM(derive_key(password, header))

    def plain():
        any_block = False
        for counter, (block, last) in enumerate(chunkify(rest(), chunk + tag_len)):
            any_block = True
            if len(block) < tag_len:
                raise SDXError("SDX 文件被截断")
            try:
                yield aes.decrypt(nonce(counter, last), block, prefix)
            except Exception:
                raise SDXError("口令不正确或文件已损坏")
        if not any_block:
            raise SDXError("SDX 文件被截断")

    pit = plain()
    buf = b""
    meta = None
    while True:
        if len(buf) >= 4:
            mlen = struct.unpack(">I", buf[:4])[0]
            if mlen > LIMITS["meta_max"]:
                raise SDXError("SDX meta too large")
            if len(buf) >= 4 + mlen:
                try:
                    meta = json.loads(buf[4:4 + mlen].decode("utf-8"))
                except Exception:
                    raise SDXError("SDX meta json")
                buf = buf[4 + mlen:]
                break
        try:
            buf += next(pit)
        except StopIteration:
            raise SDXError("SDX 文件被截断")

    def parts():
        if buf:
            yield buf
        for p in pit:
            yield p

    return ver, header, meta or {}, parts()


# ---------- 文本 ----------

def encrypt_text(plain, password):
    import io
    out = io.BytesIO()
    encrypt_stream(iter([plain.encode("utf-8")]), password, {}, out)
    return TEXT_PREFIX + base64.b64encode(out.getvalue()).decode()


def decrypt_text(s, password):
    s = s.strip()
    if s[:len(TEXT_PREFIX)].upper() == TEXT_PREFIX:
        s = s[len(TEXT_PREFIX):]
    raw = base64.b64decode("".join(s.split()))
    _, _, _, parts = open_decrypt(iter([raw]), password)
    return b"".join(parts).decode("utf-8")


# ---------- CLI ----------

def get_password(args, confirm=False):
    pw = args.password or os.environ.get("SDX_PASSWORD")
    if pw:
        return pw
    pw = getpass.getpass("口令: ")
    if confirm and getpass.getpass("再输一次: ") != pw:
        print("两次口令不一致", file=sys.stderr)
        sys.exit(2)
    if not pw:
        print("口令不能为空", file=sys.stderr)
        sys.exit(2)
    return pw


def cmd_enc(args):
    src = args.file
    out = args.output or (src + ".sdx")
    pw = get_password(args, confirm=True)
    size = os.path.getsize(src)
    meta = {"name": os.path.basename(src), "mime": "application/octet-stream", "size": size}
    with open(src, "rb") as fi, open(out, "wb") as fo:
        n = encrypt_stream(read_pieces(fi), pw, meta, fo)
    print("已加密: %s → %s (%d 字节)" % (src, out, n))


def cmd_dec(args):
    src = args.file
    pw = get_password(args)
    with open(src, "rb") as fi:
        ver, header, meta, parts = open_decrypt(read_pieces(fi), pw)
        safe = os.path.basename(str(meta.get("name") or "")).strip()
        if safe in ("", ".", ".."):
            safe = ""
        out = args.output or safe or (src[:-4] if src.lower().endswith(".sdx") else src + ".out")
        if os.path.exists(out) and not args.force:
            print("目标已存在: %s（加 -f 覆盖）" % out, file=sys.stderr)
            sys.exit(2)
        n = 0
        with open(out, "wb") as fo:
            for p in parts:
                fo.write(p)
                n += len(p)
    print("已解密 (v%d): %s → %s (%d 字节)" % (ver, src, out, n))


def cmd_info(args):
    with open(args.file, "rb") as fi:
        parsed = parse_header(fi.read(7 + MAX_HEADER))
    if not parsed:
        raise SDXError("SDX 文件过短")
    ver, header, _ = parsed
    print(json.dumps({"version": ver, "header": header}, ensure_ascii=False, indent=2))


def main(argv=None):
    ap = argparse.ArgumentParser(description="SDX 加解密命令行（文件加密 BitCipher · 星辉数盾）")
    sub = ap.add_subparsers(dest="cmd", required=True)
    for name, fn in (("enc", cmd_enc), ("dec", cmd_dec)):
        p = sub.add_parser(name)
        p.add_argument("file")
        p.add_argument("-o", "--output")
        p.add_argument("-p", "--password")
        p.add_argument("-f", "--force", action="store_true", help="覆盖已存在的输出")
        p.set_defaults(fn=fn)
    p = sub.add_parser("enc-text"); p.add_argument("text"); p.add_argument("-p", "--password")
    p.set_defaults(fn=lambda a: print(encrypt_text(a.text, get_password(a, confirm=True))))
    p = sub.add_parser("dec-text"); p.add_argument("text"); p.add_argument("-p", "--password")
    p.set_defaults(fn=lambda a: print(decrypt_text(a.text, get_password(a))))
    p = sub.add_parser("info"); p.add_argument("file")
    p.set_defaults(fn=cmd_info)
    args = ap.parse_args(argv)
    try:
        args.fn(args)
    except SDXError as e:
        print("错误: %s" % e, file=sys.stderr)
        sys.exit(1)


if __name__ == "__main__":
    main()
