new fixes
This commit is contained in:
1 parent
075a35b441
commit
ae14782da4
12 files changed
+412
-23
No files matched your search
@@ -5,7 +5,7 @@ import hashlib
|
||||
import struct
|
||||
from typing import BinaryIO, Iterator, AsyncIterator
|
||||
|
||||
from gcm_siv import GcmSiv # requires `gcm_siv` package
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCMSIV
|
||||
|
||||
|
||||
MAGIC = b"ENCF"
|
||||
@@ -40,12 +40,14 @@ def encrypt_file_to_encf(src: BinaryIO, key: bytes, chunk_bytes: int, salt: byte
|
||||
"""
|
||||
yield build_header(chunk_bytes, salt)
|
||||
idx = 0
|
||||
cipher = AESGCMSIV(key)
|
||||
|
||||
while True:
|
||||
block = src.read(chunk_bytes)
|
||||
if not block:
|
||||
break
|
||||
nonce = _derive_nonce(salt, idx)
|
||||
ct_and_tag = GcmSiv(key).encrypt(nonce, block, associated_data=None)
|
||||
ct_and_tag = cipher.encrypt(nonce, block, associated_data=None)
|
||||
# Split tag
|
||||
tag = ct_and_tag[-16:]
|
||||
ct = ct_and_tag[:-16]
|
||||
@@ -85,6 +87,8 @@ async def decrypt_encf_to_file(byte_iter: AsyncIterator[bytes], key: bytes, out_
|
||||
salt = bytes(buf[11:11 + salt_len])
|
||||
del buf[:hdr_len]
|
||||
|
||||
cipher = AESGCMSIV(key)
|
||||
|
||||
async with aiofiles.open(out_path, 'wb') as out:
|
||||
idx = 0
|
||||
TAG_LEN = 16
|
||||
@@ -103,7 +107,6 @@ async def decrypt_encf_to_file(byte_iter: AsyncIterator[bytes], key: bytes, out_
|
||||
tag = bytes(buf[p_len:p_len+TAG_LEN])
|
||||
del buf[:p_len+TAG_LEN]
|
||||
nonce = _derive_nonce(salt, idx)
|
||||
pt = GcmSiv(key).decrypt(nonce, ct + tag, associated_data=None)
|
||||
pt = cipher.decrypt(nonce, ct + tag, associated_data=None)
|
||||
await out.write(pt)
|
||||
idx += 1
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hmac
|
||||
import hashlib
|
||||
import os
|
||||
import struct
|
||||
from typing import BinaryIO, Iterator, AsyncIterator
|
||||
|
||||
import aiofiles
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
|
||||
MAGIC = b"ENCF"
|
||||
VERSION = 1
|
||||
SCHEME_AES_GCM = 0x03
|
||||
|
||||
CHUNK_BYTES = int(os.getenv("CRYPTO_CHUNK_BYTES", "1048576"))
|
||||
|
||||
|
||||
def _derive_nonce(salt: bytes, idx: int) -> bytes:
|
||||
"""Derive a deterministic 12-byte nonce from salt and chunk index."""
|
||||
if len(salt) < 12:
|
||||
raise ValueError("salt must be at least 12 bytes")
|
||||
idx_bytes = idx.to_bytes(8, "big")
|
||||
return hmac.new(salt, idx_bytes, hashlib.sha256).digest()[:12]
|
||||
|
||||
|
||||
def build_header(chunk_bytes: int, salt: bytes) -> bytes:
|
||||
if not (0 < chunk_bytes <= (1 << 31)):
|
||||
raise ValueError("chunk_bytes must be between 1 and 2^31")
|
||||
if not (1 <= len(salt) <= 255):
|
||||
raise ValueError("salt length must be 1..255 bytes")
|
||||
# MAGIC(4) | ver(1) | scheme(1) | chunk_bytes(4,BE) | salt_len(1) | salt | reserved(5 zeros)
|
||||
hdr = bytearray()
|
||||
hdr += MAGIC
|
||||
hdr.append(VERSION)
|
||||
hdr.append(SCHEME_AES_GCM)
|
||||
hdr += struct.pack(">I", int(chunk_bytes))
|
||||
hdr.append(len(salt))
|
||||
hdr += salt
|
||||
hdr += b"\x00" * 5
|
||||
return bytes(hdr)
|
||||
|
||||
|
||||
def encrypt_file_to_encf(src: BinaryIO, key: bytes, chunk_bytes: int, salt: bytes) -> Iterator[bytes]:
|
||||
"""Yield ENCF v1 frames encrypted with AES-GCM."""
|
||||
if len(key) not in (16, 24, 32):
|
||||
raise ValueError("AES-GCM key must be 128, 192 or 256 bits long")
|
||||
cipher = AESGCM(key)
|
||||
yield build_header(chunk_bytes, salt)
|
||||
idx = 0
|
||||
while True:
|
||||
block = src.read(chunk_bytes)
|
||||
if not block:
|
||||
break
|
||||
nonce = _derive_nonce(salt, idx)
|
||||
ct = cipher.encrypt(nonce, block, associated_data=None)
|
||||
tag = ct[-16:]
|
||||
data = ct[:-16]
|
||||
yield struct.pack(">I", len(block))
|
||||
yield data
|
||||
yield tag
|
||||
idx += 1
|
||||
|
||||
|
||||
async def decrypt_encf_to_file(byte_iter: AsyncIterator[bytes], key: bytes, out_path: str) -> None:
|
||||
"""Parse ENCF v1 (AES-GCM) stream and write plaintext to `out_path`."""
|
||||
if len(key) not in (16, 24, 32):
|
||||
raise ValueError("AES-GCM key must be 128, 192 or 256 bits long")
|
||||
cipher = AESGCM(key)
|
||||
buf = bytearray()
|
||||
|
||||
async def _fill(n: int) -> None:
|
||||
nonlocal buf
|
||||
while len(buf) < n:
|
||||
try:
|
||||
chunk = await byte_iter.__anext__()
|
||||
except StopAsyncIteration:
|
||||
break
|
||||
if chunk:
|
||||
buf.extend(chunk)
|
||||
|
||||
# Parse header
|
||||
await _fill(11)
|
||||
if buf[:4] != MAGIC:
|
||||
raise ValueError("bad magic")
|
||||
version = buf[4]
|
||||
scheme = buf[5]
|
||||
if version != VERSION or scheme != SCHEME_AES_GCM:
|
||||
raise ValueError("unsupported ENCF header")
|
||||
chunk_bytes = struct.unpack(">I", bytes(buf[6:10]))[0]
|
||||
salt_len = buf[10]
|
||||
hdr_len = 4 + 1 + 1 + 4 + 1 + salt_len + 5
|
||||
await _fill(hdr_len)
|
||||
salt = bytes(buf[11:11 + salt_len])
|
||||
del buf[:hdr_len]
|
||||
|
||||
async with aiofiles.open(out_path, "wb") as out:
|
||||
idx = 0
|
||||
TAG_LEN = 16
|
||||
while True:
|
||||
await _fill(4)
|
||||
if len(buf) == 0:
|
||||
break
|
||||
if len(buf) < 4:
|
||||
raise ValueError("truncated frame length")
|
||||
p_len = struct.unpack(">I", bytes(buf[:4]))[0]
|
||||
del buf[:4]
|
||||
await _fill(p_len + TAG_LEN)
|
||||
if len(buf) < p_len + TAG_LEN:
|
||||
raise ValueError("truncated cipher/tag")
|
||||
ct = bytes(buf[:p_len])
|
||||
tag = bytes(buf[p_len:p_len + TAG_LEN])
|
||||
del buf[:p_len + TAG_LEN]
|
||||
nonce = _derive_nonce(salt, idx)
|
||||
pt = cipher.decrypt(nonce, ct + tag, associated_data=None)
|
||||
await out.write(pt)
|
||||
idx += 1
|
||||
@@ -0,0 +1,120 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional
|
||||
|
||||
from .aes_gcm_stream import CHUNK_BYTES, encrypt_file_to_encf
|
||||
from .encf_stream import decrypt_encf_auto
|
||||
from .keywrap import unwrap_dek, KeyWrapError
|
||||
|
||||
|
||||
def _normalize_base64(value: str) -> str:
|
||||
padding = (-len(value)) % 4
|
||||
if padding:
|
||||
return value + "=" * padding
|
||||
return value
|
||||
|
||||
|
||||
def _decode_key(value: str, fmt: str) -> bytes:
|
||||
if fmt == "base64":
|
||||
return base64.b64decode(_normalize_base64(value))
|
||||
if fmt == "hex":
|
||||
cleaned = value[2:] if value.lower().startswith("0x") else value
|
||||
return bytes.fromhex(cleaned)
|
||||
if fmt == "raw":
|
||||
return value.encode()
|
||||
raise ValueError(f"unsupported key format: {fmt}")
|
||||
|
||||
|
||||
def _decode_salt(value: str, fmt: str) -> bytes:
|
||||
if fmt == "base64":
|
||||
return base64.b64decode(_normalize_base64(value))
|
||||
if fmt == "hex":
|
||||
cleaned = value[2:] if value.lower().startswith("0x") else value
|
||||
return bytes.fromhex(cleaned)
|
||||
raise ValueError(f"unsupported salt format: {fmt}")
|
||||
|
||||
|
||||
async def _decrypt_file(input_path: str, key: bytes, output_path: str) -> None:
|
||||
async def _aiter():
|
||||
with open(input_path, "rb") as src:
|
||||
while True:
|
||||
chunk = src.read(65536)
|
||||
if not chunk:
|
||||
break
|
||||
yield chunk
|
||||
|
||||
await decrypt_encf_auto(_aiter(), key, output_path)
|
||||
|
||||
|
||||
def cmd_encrypt(args: argparse.Namespace) -> int:
|
||||
key = _decode_key(args.key, args.key_format)
|
||||
salt = _decode_salt(args.salt, args.salt_format) if args.salt else os.urandom(args.salt_bytes)
|
||||
os.makedirs(os.path.dirname(args.output) or ".", exist_ok=True)
|
||||
with open(args.input, "rb") as src, open(args.output, "wb") as dst:
|
||||
for chunk in encrypt_file_to_encf(src, key, args.chunk_bytes, salt):
|
||||
dst.write(chunk)
|
||||
# Emit JSON metadata with salt for convenience
|
||||
meta = {
|
||||
"salt_b64": base64.b64encode(salt).decode(),
|
||||
"chunk_bytes": args.chunk_bytes,
|
||||
"aead_scheme": "AES_GCM",
|
||||
}
|
||||
print(json.dumps(meta), file=sys.stdout)
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_decrypt(args: argparse.Namespace) -> int:
|
||||
if bool(args.key) == bool(args.wrapped_key):
|
||||
raise SystemExit("Provide exactly one of --key or --wrapped-key")
|
||||
if args.wrapped_key:
|
||||
try:
|
||||
key = unwrap_dek(args.wrapped_key)
|
||||
except KeyWrapError as exc:
|
||||
raise SystemExit(f"Failed to unwrap key: {exc}") from exc
|
||||
else:
|
||||
key = _decode_key(args.key, args.key_format)
|
||||
os.makedirs(os.path.dirname(args.output) or ".", exist_ok=True)
|
||||
asyncio.run(_decrypt_file(args.input, key, args.output))
|
||||
return 0
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(prog="python -m app.core.crypto.cli", description="ENCF AES-GCM helper")
|
||||
sub = parser.add_subparsers(dest="command", required=True)
|
||||
|
||||
enc = sub.add_parser("encrypt", help="Encrypt file into ENCF v1 stream (AES-256-GCM)")
|
||||
enc.add_argument("--input", required=True, help="Path to plaintext input file")
|
||||
enc.add_argument("--output", required=True, help="Destination path for ENCF output")
|
||||
enc.add_argument("--key", required=True, help="Encryption key")
|
||||
enc.add_argument("--key-format", choices=["base64", "hex", "raw"], default="base64")
|
||||
enc.add_argument("--salt", help="Salt in specified format; generates random if omitted")
|
||||
enc.add_argument("--salt-format", choices=["base64", "hex"], default="base64")
|
||||
enc.add_argument("--salt-bytes", type=int, default=16, help="Salt length when generated (default: 16)")
|
||||
enc.add_argument("--chunk-bytes", type=int, default=CHUNK_BYTES, help="Plaintext chunk size (default from env)")
|
||||
enc.set_defaults(func=cmd_encrypt)
|
||||
|
||||
dec = sub.add_parser("decrypt", help="Decrypt ENCF stream to plaintext")
|
||||
dec.add_argument("--input", required=True, help="Path to ENCF input file")
|
||||
dec.add_argument("--output", required=True, help="Destination path for decrypted file")
|
||||
dec.add_argument("--key", help="Plaintext key")
|
||||
dec.add_argument("--wrapped-key", help="Wrapped key produced by the backend")
|
||||
dec.add_argument("--key-format", choices=["base64", "hex", "raw"], default="base64")
|
||||
dec.set_defaults(func=cmd_decrypt)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: Optional[list[str]] = None) -> int:
|
||||
parser = build_parser()
|
||||
args = parser.parse_args(argv)
|
||||
return args.func(args)
|
||||
|
||||
|
||||
if __name__ == "__main__": # pragma: no cover
|
||||
sys.exit(main())
|
||||
@@ -4,6 +4,7 @@ from typing import AsyncIterator
|
||||
|
||||
from .aes_gcm_siv_stream import MAGIC as _MAGIC, VERSION as _VER, SCHEME_AES_GCM_SIV
|
||||
from .aes_gcm_siv_stream import decrypt_encf_to_file as _dec_gcmsiv
|
||||
from .aes_gcm_stream import SCHEME_AES_GCM, decrypt_encf_to_file as _dec_gcm
|
||||
from .aes_siv_stream import decrypt_encf_to_file as _dec_siv
|
||||
|
||||
|
||||
@@ -38,6 +39,7 @@ async def decrypt_encf_auto(byte_iter: AsyncIterator[bytes], key: bytes, out_pat
|
||||
|
||||
if scheme == SCHEME_AES_GCM_SIV:
|
||||
await _dec_gcmsiv(_prepend_iter(), key, out_path)
|
||||
elif scheme == SCHEME_AES_GCM:
|
||||
await _dec_gcm(_prepend_iter(), key, out_path)
|
||||
else:
|
||||
await _dec_siv(_prepend_iter(), key, out_path)
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import os
|
||||
import threading
|
||||
from typing import Optional
|
||||
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
|
||||
_VERSION = 1
|
||||
_PREFIX_LEN = 1 # version byte
|
||||
_NONCE_LEN = 12
|
||||
_TAG_LEN = 16
|
||||
_valid_key_lengths = {16, 24, 32}
|
||||
_kek_lock = threading.Lock()
|
||||
_cached_kek: Optional[bytes] = None
|
||||
|
||||
|
||||
class KeyWrapError(RuntimeError):
|
||||
"""Raised when KEK configuration or unwrap operations fail."""
|
||||
|
||||
|
||||
def _normalize_base64(value: str) -> str:
|
||||
v = value.strip()
|
||||
missing = (-len(v)) % 4
|
||||
if missing:
|
||||
v += "=" * missing
|
||||
return v
|
||||
|
||||
|
||||
def _decode_key_material(value: str) -> bytes:
|
||||
v = value.strip()
|
||||
if v.startswith("0x") or v.startswith("0X"):
|
||||
v = v[2:]
|
||||
try:
|
||||
raw = bytes.fromhex(v)
|
||||
if len(raw) in _valid_key_lengths:
|
||||
return raw
|
||||
except ValueError:
|
||||
pass
|
||||
try:
|
||||
raw = base64.b64decode(_normalize_base64(value), validate=False)
|
||||
if len(raw) in _valid_key_lengths:
|
||||
return raw
|
||||
except Exception as exc: # noqa: BLE001 - we want to re-raise as KeyWrapError
|
||||
raise KeyWrapError(f"invalid KEK encoding: {exc}") from exc
|
||||
raise KeyWrapError("KEK must decode to 16/24/32 bytes")
|
||||
|
||||
|
||||
def _load_kek() -> bytes:
|
||||
global _cached_kek
|
||||
if _cached_kek is not None:
|
||||
return _cached_kek
|
||||
with _kek_lock:
|
||||
if _cached_kek is not None:
|
||||
return _cached_kek
|
||||
env = os.getenv("CONTENT_KEY_KEK_B64") or os.getenv("CONTENT_KEY_KEK_HEX")
|
||||
if not env:
|
||||
raise KeyWrapError("CONTENT_KEY_KEK_B64 or CONTENT_KEY_KEK_HEX must be set")
|
||||
kek = _decode_key_material(env)
|
||||
if len(kek) != 32:
|
||||
# Force 256-bit KEK for uniform security properties
|
||||
raise KeyWrapError("KEK must be 32 bytes (256-bit) for AES-256-GCM")
|
||||
_cached_kek = kek
|
||||
return _cached_kek
|
||||
|
||||
|
||||
def wrap_dek(plaintext: bytes) -> str:
|
||||
"""Wrap a DEK (plaintext bytes) with AES-256-GCM; return base64 string."""
|
||||
if not isinstance(plaintext, (bytes, bytearray)):
|
||||
raise TypeError("plaintext must be bytes")
|
||||
kek = _load_kek()
|
||||
nonce = os.urandom(_NONCE_LEN)
|
||||
cipher = AESGCM(kek)
|
||||
ct = cipher.encrypt(nonce, bytes(plaintext), associated_data=None)
|
||||
blob = bytes([_VERSION]) + nonce + ct
|
||||
return base64.b64encode(blob).decode()
|
||||
|
||||
|
||||
def unwrap_dek(encoded: str) -> bytes:
|
||||
"""Unwrap DEK from base64 string. Supports legacy (raw base64 key) values."""
|
||||
if not encoded:
|
||||
raise KeyWrapError("empty key payload")
|
||||
try:
|
||||
raw = base64.b64decode(_normalize_base64(encoded), validate=False)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise KeyWrapError(f"invalid base64 payload: {exc}") from exc
|
||||
if not raw:
|
||||
raise KeyWrapError("decoded payload is empty")
|
||||
version = raw[0]
|
||||
if version == _VERSION:
|
||||
if len(raw) < _PREFIX_LEN + _NONCE_LEN + _TAG_LEN + 1:
|
||||
raise KeyWrapError("wrapped payload too short")
|
||||
nonce = raw[_PREFIX_LEN:_PREFIX_LEN + _NONCE_LEN]
|
||||
ciphertext = raw[_PREFIX_LEN + _NONCE_LEN:]
|
||||
kek = _load_kek()
|
||||
cipher = AESGCM(kek)
|
||||
try:
|
||||
return cipher.decrypt(nonce, ciphertext, associated_data=None)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
raise KeyWrapError(f"unwrap failed: {exc}") from exc
|
||||
# Legacy fallback: value is raw DEK (no version prefix)
|
||||
if len(raw) in {16, 24, 32}:
|
||||
return raw
|
||||
raise KeyWrapError("unknown key payload format")
|
||||
Reference in new issue
Block a user