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, "", "eval"), {}) def run_module(source: str) -> str: stdout = io.StringIO() with contextlib.redirect_stdout(stdout): exec(compile(source, "", "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, "", "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, "", "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()