Coverage for src/renaissance/recipes/unit_to_pytest.py: 97%
173 statements
« prev ^ index » next coverage.py v7.16.1, created at 2026-10-02 13:12 +0000
« prev ^ index » next coverage.py v7.16.1, created at 2026-10-02 13:12 +0000
1"""AI: Recipe that converts unittest-style test files to pytest style."""
3import textwrap
4from collections.abc import Sequence
5from pathlib import Path
7from renaissance.integrations.python.ast.util import convert_function
8from renaissance.recipes.python_refactoring import PythonRefactoring
9from renaissance.syntax_tree import PatternMatch
10from renaissance.syntax_tree.ast_finder import find_semantic_kind
11from renaissance.syntax_tree.match_finder import match_pattern
12from renaissance.syntax_tree.node_protocol import NodeProtocol
13from renaissance.syntax_tree.semantic_kind import SemanticKind
16class UnitToPytest(PythonRefactoring):
17 """AI: Recipe that converts unittest-style test files to pytest style."""
19 def __init__(self, file):
20 """Hide internal administration in the parent class so that this class you only deals with specific refactors."""
21 super().__init__(file)
22 self.black_list_pattern = "utils_for_test"
23 self.white_list_pattern = "test"
25 def run(self):
26 """Entry point for converting unittest to pytest."""
27 self.refactor()
29 self.post_processing()
31 def refactor(self):
32 """AI: Apply file-, class-, and function-level unittest-to-pytest conversions."""
33 # 1: file level changes
34 self.convert_test_class()
35 self.restructure_module()
36 self.replace_stmt("unittest.main()", "pytest.main()")
37 self.replace_stmt("import unittest", "import pytest\nfrom hamcrest import *")
38 self.replace_stmt("from parameterized import parameterized", "import pytest\nfrom hamcrest import *")
39 self.replace_stmt("from unittest import TestCase,$$symbols", "import pytest\nfrom hamcrest import *")
40 self.replace_stmt("from unittest import TestCase", "import pytest\nfrom hamcrest import *")
42 # 2: class level changes
43 self.convert_parameterized_test()
44 self.convert_test_setup()
45 self.commit()
47 # 3: function level changes
49 self.convert_skip_test()
50 self.remove_print()
51 self.convert_plain_assert_same_length()
52 self.commit()
53 self.replace_stmt("assert $stmt, $$msg", "assert_that($stmt, is_(True), $$msg)")
54 self.replace_stmt("self.assertTrue($exp,$$msg)", "assert_that($exp, is_(True), $$msg)")
55 self.replace_stmt("self.assertFalse($exp, $$msg)", "assert_that($exp, is_(False), $$msg)")
57 self.convert_assert("self.assertEqual($exp, $act)", "assert_that($exp, is_($act))")
58 self.convert_assert("self.assertGreaterEqual($exp, $act)", "assert_that($exp, greater_than_or_equal_to($act))")
59 self.convert_assert("self.assertGreater($exp, $act)", "assert_that($exp, greater_than($act))")
60 self.convert_assert("self.assertLesserEqual($exp, $act)", "assert_that($exp, less_than_or_equal_to($act))")
61 self.convert_assert("self.assertLesser($exp, $act)", "assert_that($exp, less_than($act))")
62 self.convert_assert("self.assertMultiLineEqual($act, $exp)", "assert_that($act, is_($exp))")
64 self.replace_stmt("self.assertIn($act, $exp)", "assert_that($exp, contain_string($act))")
65 self.replace_stmt("self.assertIsInstance($act, $exp)", "assert_that($act, is_($exp))")
66 self.replace_stmt("with self.assertRaises($exc): $call()", "assert_that(calling($call), raises($exc))")
68 def post_processing(self):
69 """AI: Repeatedly simplify assert_that(...) expressions until no further changes occur."""
70 # 4: improve to more concise asserts
71 while self.has_changed():
72 self.commit()
73 self.replace_stmt("assert_that($exp)", "assert_that($exp, is_(True))")
74 self.replace_stmt("assert_that(isinstance($exp, $act))", "assert_that($exp, is_($act))")
75 self.replace_stmt("assert_that(len($exp), $act)", "assert_that($exp, has_length($act))")
76 self.replace_stmt("assert_that(len($exp) >= 1)", "assert_that($exp, is_not(empty()))")
77 self.replace_stmt("assert_that(len($exp) >= 1, is_(True))", "assert_that($exp, is_not(empty()))")
78 self.replace_stmt("assert_that(len($exp) == $length)", "assert_that($exp, has_length($length))")
79 self.replace_stmt("assert_that($exp == $act)", "assert_that($exp, is_($act), $$msg)")
80 self.replace_stmt("assert_that($exp == $act, is_(True), $$msg)", "assert_that($exp, is_($act), $$msg)")
81 self.replace_stmt("assert_that(not $stmt, is_(True), $$msg)", "assert_that($stmt, is_(False) ,$$msg)")
82 self.replace_stmt("assert_that($stmt, is_not(True), $$msg)", "assert_that($stmt, is_(False) ,$$msg)")
83 self.replace_stmt("assert_that(not $stmt)", "assert_that($stmt, is_(False))")
84 self.replace_stmt("assert_that($el in $col, is_(True))", "assert_that($col, contains_exactly($el))")
85 self.replace_stmt("assert_that($exp, has_length(is_($act)))", "assert_that($exp, has_length($act))")
86 self.swap_expected_and_actual()
87 self.replace_stmt("assert_that(not $stmt)", "assert_that($stmt, is_(False))")
88 self.replace_stmt("assert_that($exp.startswith($act))", "assert_that($exp, starts_with($act))")
89 self.remove_duplicate_import("import pytest\nfrom hamcrest import *")
90 self.commit()
92 def convert_test_class(self):
93 """AI: Rewrite TestCase-derived class headers to drop the unittest base class."""
94 test_main: Sequence[NodeProtocol] = self.pattern_factory.create_statements(
95 "class $klass($test_class):\n $$test_cases\n",
96 ) # type: ignore[assignment]
97 for match in match_pattern(self.root.children, test_main):
98 klass = match["$klass"]
99 test_class = match["$test_class"]
101 if test_class.endswith("TestCase"):
102 # class inherit from TestCase (or unittest.TestCase)
103 if klass.endswith("Test"):
104 # class name ends with Test, rename by move Test to front
105 repl = match.signature.replace(f"{klass}({test_class}):", f"Test{klass[:-4]}:")
106 else:
107 # we assume there are only 2 variant TestExample and ExampleTest
108 repl = match.signature.replace(f"({test_class}):", ":")
110 # repl = f'class {match.expansions["$klass"][0]}:\n{raw(match.expansions["$$test_cases"])}'
111 self.replace(repl, match.nodes, False, False)
113 def convert_test_setup(self):
114 """AI: Convert a setUp method into a pytest autouse fixture named setup."""
115 setup_function = self.pattern_factory.create_statements("def setUp(self): $$stmts")
116 for match in match_pattern(self.body, setup_function):
117 # add decorator to the setup dunction and convert to snake case
118 repl = f"@pytest.fixture(autouse=True)\n{match.signature}".replace(" setUp(self)", " setup(self)")
119 self.replace(repl, match.nodes, False, False)
121 def convert_assert(self, pattern, replacement):
122 """AI: Replace calls matching pattern with replacement, swapping expected/actual arguments as needed."""
123 pat = self.pattern_factory.create_statements(pattern)
124 for match in match_pattern(self.root.children, pat):
125 repl = replacement
126 if self.is_swapped(match):
127 exp = match["$act"]
128 act = match["$exp"]
129 else: # original is wrong
130 act = match["$act"]
131 exp = match["$exp"]
132 repl = repl.replace("$exp", exp).replace("$act", act)
133 self.replace(repl, match.nodes, False, False)
135 def is_swapped(self, match: PatternMatch) -> bool:
136 """AI: Return whether $exp and $act appear swapped in the match (i.e. $exp is a literal)."""
137 return match.expansions["$exp"][0].semantic_kind is SemanticKind.LITERAL
139 def convert_parameterized_test(self):
140 """AI: Convert @parameterized.expand-decorated test functions into @pytest.mark.parametrize."""
141 unittest = self.pattern_factory.create_statements(
142 textwrap.dedent("""
143 @parameterized.expand($$parameters)
144 @$$decorator
145 def $fun($$args, *$$varg):
146 $$stmts
147 """),
148 )
149 for match in match_pattern(self.root.children, unittest):
150 fun = match.nodes[0]
151 args = ", ".join([arg.node.arg for arg in match.expansions["$$args"]])
152 if varg := match.expansions["$$varg"]:
153 args = f"{args}, *{varg[0].signature}"
154 args = args.replace("self, ", "")
155 repl = fun.signature
156 if " def " in repl:
157 repl = repl.replace("@parameterized.expand(", f' @pytest.mark.parametrize("{args}",')
158 repl = repl.replace("@unittest.skip(", "@pytest.mark.skip(")
159 repl = textwrap.dedent(repl)
160 else:
161 repl = repl.replace("@parameterized.expand(", f'@pytest.mark.parametrize("{args}",')
162 repl = repl.replace("@unittest.skip(", "@pytest.mark.skip(")
164 self.replace(repl, fun, False, False)
166 def remove_print(self):
167 """AI: Remove print(...) statements, or their containing block if it's the only statement."""
168 print_msg = self.pattern_factory.create_statements("print($$msg)") # type: ignore[assignment]
169 for match in match_pattern(self.root.children, print_msg):
170 if len(match.nodes[0].parent.parent.body) == 1:
171 self.remove([match.nodes[0].parent.parent], False, False)
172 else:
173 self.remove(match.nodes, False, False)
175 def convert_plain_assert_same_length(self):
176 """AI: Replace a manual length-check assert with an assert_that(...) has_length assertion."""
177 pattern: Sequence[NodeProtocol] = self.pattern_factory.create_statements(
178 '$act: int = len($real)\nassert $exp == $act, "$act = " + str($act)',
179 )
180 for match in match_pattern(self.body, pattern):
181 repl = 'assert_that($real, has_length($exp), f"length of $real = {len($real)}")'
182 real = match["$real"]
183 exp = match["$exp" if self.is_swapped(match) else "$act"] # use "$act" when original is wrong
184 repl = repl.replace("$exp", exp).replace("$real", real)
185 self.replace(repl, match.nodes, False, False)
187 def convert_skip_test(self):
188 """AI: Replace unittest.skip attribute references with pytest.mark.skip."""
189 nodes = find_semantic_kind(self.root, SemanticKind.ATTRIBUTE)
190 for node in nodes:
191 if node.signature == "unittest.skip":
192 self.replace("pytest.mark.skip", node, False, False)
194 def swap_expected_and_actual(self):
195 """AI: Swap the $exp and $act arguments of assert_that(...) calls when they appear reversed."""
196 pattern: Sequence[NodeProtocol] = self.pattern_factory.create_statements("assert_that($exp, is_($act))") # type: ignore[assignment]
197 for match in match_pattern(self.root.children, pattern):
198 if self.is_swapped(match):
199 repl = "assert_that($act, is_($exp))"
200 act = match["$act"]
201 exp = match["$exp"]
202 repl = repl.replace("$exp", exp).replace("$act", act)
203 self.replace(repl, match.nodes, False, False)
205 def restructure_module(self):
206 """AI: Move module-level functions into a test class, creating one if none exists."""
207 funs = [stmt for stmt in self.body if stmt.semantic_kind is SemanticKind.FUNCTION]
208 test_classes = [stmt for stmt in self.body if stmt.semantic_kind is SemanticKind.CLASS and stmt.name.startswith("Test")]
209 if len(funs) == 0:
210 return
211 if len(test_classes) == 0:
212 # file does not contain any test class, create a new class and add function in class
213 cls = f"class {self.convert_file_to_test_class()}:\n"
214 for fun in funs:
215 cls += textwrap.indent(convert_function(fun), " ")
216 self.remove([fun])
217 self.insert_before(cls, funs[0])
218 else:
219 # one or more class in file, add function as member of the last class in file
220 for fun in funs:
221 # assuming the class comes first
222 meth = convert_function(fun)
223 self.insert_after(meth, test_classes[-1].body[-1])
224 self.remove(fun)
225 self.commit()
226 for fun in funs:
227 # also change the calling signature of those functions in case they are not test cases
228 function_call = [self.pattern_factory.create_expression(f"{fun.name}($$args)")]
229 for call in match_pattern(self.root.children, function_call):
230 sig = call.nodes[0].signature
231 self.replace(f"self.{sig}", call.nodes, False, False)
232 self.commit()
234 def convert_file_to_test_class(self):
235 """AI: Derive a PascalCase Test-prefixed class name from the file's stem."""
236 path = Path(self.filename)
237 stem = path.stem
238 parts = stem.split("_")
239 if parts[-1].lower() == "test":
240 parts = parts[:-1]
241 name = "".join(word.capitalize() for word in parts)
242 return name if name.startswith("Test") else f"Test{name}"
244 def remove_duplicate_import(self, import_str):
245 """AI: Remove duplicate occurrences of the given import statement, keeping the first and last."""
246 import_stmt: Sequence[NodeProtocol] = self.pattern_factory.create_statements(import_str) # type: ignore[assignment]
247 # type: ignore[assignment]
248 duplicate_imports = match_pattern(self.body, import_stmt)
250 for match in duplicate_imports[1:-1]:
251 self.remove(match.nodes, False, False)