Files
bikini-patchwork/tests/test_transforms.py
T
2026-06-17 20:45:45 -05:00

101 lines
3.7 KiB
Python

from __future__ import annotations
import ast
import contextlib
import io
import random
import sys
import unittest
from pathlib import Path
ROOT = Path(__file__).resolve().parent.parent
sys.path.insert(0, str(ROOT))
from patchwork.transforms.identifiers import IdentifierRenamer
from patchwork.transforms.normalize import MatchLowerer
from patchwork.transforms.numbers import NumberObfuscator
from patchwork.transforms.strings import StringEncryptor, build_decrypt_helper
def eval_ast_expr(node: ast.AST) -> object:
expr = ast.Expression(body=node)
ast.fix_missing_locations(expr)
return eval(compile(expr, "<test>", "eval"), {})
def run_module(source: str) -> str:
stdout = io.StringIO()
with contextlib.redirect_stdout(stdout):
exec(compile(source, "<test>", "exec"), {})
return stdout.getvalue()
class TransformTests(unittest.TestCase):
def test_new_number_transforms_preserve_values(self) -> None:
values = [-65537, -129, -1, 0, 1, 42, 4096, 2**40 + 123]
methods = ["_t_invert", "_t_bitmask", "_t_divmod", "_t_table"]
for method_name in methods:
for value in values:
with self.subTest(method=method_name, value=value):
obfuscator = NumberObfuscator(random.Random(2026), prob=1.0)
node = getattr(obfuscator, method_name)(value)
self.assertEqual(value, eval_ast_expr(node))
def test_string_encryptor_builds_shared_literal_pool(self) -> None:
source = (
"VALUE = 'PATCHWORK_SECRET_MARKER_2026'\n"
"RAW = b'PATCHWORK_BYTES_MARKER_2026'\n"
"print(VALUE, RAW.decode())\n"
)
tree = ast.parse(source)
encryptor = StringEncryptor(random.Random(17), "_pw_dec_str", "_pw_dec_bytes", "_pw_literal_pool")
tree = encryptor.visit(tree)
pool_stmt = encryptor.build_pool_stmt()
self.assertIsNotNone(pool_stmt)
tree.body = [pool_stmt, *build_decrypt_helper("_pw_dec_str", "_pw_dec_bytes"), *tree.body]
ast.fix_missing_locations(tree)
transformed_source = ast.unparse(tree)
self.assertIn("_pw_literal_pool", transformed_source)
self.assertNotIn("PATCHWORK_SECRET_MARKER_2026", transformed_source)
self.assertNotIn("PATCHWORK_BYTES_MARKER_2026", transformed_source)
stdout = io.StringIO()
with contextlib.redirect_stdout(stdout):
exec(compile(tree, "<test>", "exec"), {})
self.assertEqual("PATCHWORK_SECRET_MARKER_2026 PATCHWORK_BYTES_MARKER_2026\n", stdout.getvalue())
def test_reserved_import_keeps_original_binding(self) -> None:
tree = ast.parse("from collections import deque\nprint(deque([1, 2, 3]))\n")
renamed = IdentifierRenamer(random.Random(5), keep={"deque"}).visit(tree)
ast.fix_missing_locations(renamed)
source = ast.unparse(renamed)
self.assertIn("from collections import deque", source)
stdout = io.StringIO()
with contextlib.redirect_stdout(stdout):
exec(compile(renamed, "<test>", "exec"), {})
self.assertIn("deque([1, 2, 3])", stdout.getvalue())
def test_match_lowerer_preserves_general_sequence_patterns(self) -> None:
source = """
from collections import deque
def classify(value):
match value:
case (first, *middle, last):
return f"{type(value).__name__}:{first}:{middle}:{last}"
case _:
return f"{type(value).__name__}:no"
print(classify(range(4)))
print(classify("abcd"))
print(classify(deque([1, 2, 3])))
"""
tree = MatchLowerer().visit(ast.parse(source))
ast.fix_missing_locations(tree)
self.assertEqual(run_module(source), run_module(ast.unparse(tree)))
if __name__ == "__main__":
unittest.main()