Files

146 lines
5.3 KiB
Python

"""crypto 模块行为测试 (真实加密/解密/篡改, 非 mock)
运行: cd /home/lou/文档/7z-encrypt && .venv/bin/python -m unittest discover -v
"""
import io
import os
import sys
import tempfile
import unittest
from pathlib import Path
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from crypto import CryptoEngine, CryptoError, IntegrityError, KeyNotFoundError
class TestKeyring(unittest.TestCase):
def setUp(self):
self._td = tempfile.TemporaryDirectory(prefix='hermes-test-')
self.addCleanup(self._td.cleanup)
self.kr = os.path.join(self._td.name, 'keyring.json')
self.eng = CryptoEngine(self.kr)
def test_encrypt_name_roundtrip(self):
# 文件名加密回环: encrypt -> decrypt == 原名, 带 enc: 前缀, URL-safe (无 / + =)
engine = CryptoEngine(self.kr)
engine.generate_key("default_key")
token = engine.encrypt_name("机密文件.apk")
self.assertTrue(token.startswith("enc:"))
self.assertNotIn("机密文件", token) # 明文不出现在密文里
body = token[4:]
self.assertNotIn("/", body) # URL-safe: 不含路径分隔符和 + (可能含尾部 = 填充)
self.assertEqual(engine.decrypt_name(token), "机密文件.apk")
def test_encrypt_name_wrong_key_fails(self):
engine = CryptoEngine(self.kr)
engine.generate_key("default_key")
token = engine.encrypt_name("a.bin")
engine2 = CryptoEngine(self.kr + "2")
engine2.generate_key("default_key")
with self.assertRaises(IntegrityError):
engine2.decrypt_name(token)
def test_default_keyring_env_override(self):
# SZ_KEYRING 环境变量指定密钥路径 (系统级密钥位置); XDG 默认兜底
import crypto as crypto_mod
old = os.environ.get("SZ_KEYRING")
old_xdg = os.environ.get("XDG_DATA_HOME")
try:
os.environ["SZ_KEYRING"] = "/tmp/sz-kr-test/kr.json"
self.assertEqual(crypto_mod._default_keyring(), Path("/tmp/sz-kr-test/kr.json"))
# 未设 env 时用 XDG 数据目录
os.environ.pop("SZ_KEYRING", None)
os.environ["XDG_DATA_HOME"] = "/tmp/sz-xdg-test"
self.assertEqual(crypto_mod._default_keyring(), Path("/tmp/sz-xdg-test/7z-encrypt/keyring.json"))
# 都没有 -> ~/.local/share
os.environ.pop("XDG_DATA_HOME", None)
self.assertEqual(
crypto_mod._default_keyring(),
Path.home() / ".local" / "share" / "7z-encrypt" / "keyring.json",
)
finally:
if old:
os.environ["SZ_KEYRING"] = old
else:
os.environ.pop("SZ_KEYRING", None)
if old_xdg:
os.environ["XDG_DATA_HOME"] = old_xdg
else:
os.environ.pop("XDG_DATA_HOME", None)
def test_generate_and_persist(self):
k1 = self.eng.generate_key('k')
self.assertEqual(len(k1), 32)
self.assertTrue(os.path.exists(self.kr))
self.assertEqual(oct(os.stat(self.kr).st_mode & 0o777), '0o600')
# 重新实例化读回同一密钥 (持久化)
eng2 = CryptoEngine(self.kr)
self.assertEqual(eng2.get_key('k'), k1)
def test_no_overwrite(self):
self.eng.generate_key('k')
with self.assertRaises(CryptoError):
self.eng.generate_key('k')
def test_unknown_key(self):
with self.assertRaises(KeyNotFoundError):
self.eng.get_key('nope')
class TestRoundTrip(unittest.TestCase):
def setUp(self):
self._td = tempfile.TemporaryDirectory(prefix='hermes-test-')
self.addCleanup(self._td.cleanup)
self.eng = CryptoEngine(os.path.join(self._td.name, 'keyring.json'))
self.eng.generate_key('k')
def _roundtrip(self, data: bytes):
enc = io.BytesIO()
params = self.eng.encrypt_stream(io.BytesIO(data), enc, 'k')
enc.seek(0)
dec = io.BytesIO()
self.eng.decrypt_stream(enc, dec, 'k')
self.assertEqual(dec.getvalue(), data)
return params
def test_empty(self):
self._roundtrip(b'')
def test_small(self):
self._roundtrip(b'hello')
def test_cross_block(self):
# 跨 16KiB 块边界
self._roundtrip(b'x' * 16385)
def test_multi_mb_random(self):
self._roundtrip(os.urandom(5 << 20))
def test_params(self):
params = self._roundtrip(b'z')
self.assertEqual(params['alg'], 'cobblestone-aes256gcm')
self.assertEqual(params['key_id'], 'k')
self.assertEqual(params['context'], '7z-encrypt:v1')
class TestTamper(unittest.TestCase):
def setUp(self):
self._td = tempfile.TemporaryDirectory(prefix='hermes-test-')
self.addCleanup(self._td.cleanup)
self.eng = CryptoEngine(os.path.join(self._td.name, 'keyring.json'))
self.eng.generate_key('k')
def test_tamper_detected(self):
data = b'y' * 100000
enc = io.BytesIO()
self.eng.encrypt_stream(io.BytesIO(data), enc, 'k')
bad = bytearray(enc.getvalue())
bad[len(bad) // 2] ^= 0xFF
with self.assertRaises(IntegrityError):
self.eng.decrypt_stream(io.BytesIO(bytes(bad)), io.BytesIO(), 'k')
if __name__ == '__main__':
unittest.main()