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

1"""AI: Recipe that converts unittest-style test files to pytest style.""" 

2 

3import textwrap 

4from collections.abc import Sequence 

5from pathlib import Path 

6 

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 

14 

15 

16class UnitToPytest(PythonRefactoring): 

17 """AI: Recipe that converts unittest-style test files to pytest style.""" 

18 

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" 

24 

25 def run(self): 

26 """Entry point for converting unittest to pytest.""" 

27 self.refactor() 

28 

29 self.post_processing() 

30 

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 *") 

41 

42 # 2: class level changes 

43 self.convert_parameterized_test() 

44 self.convert_test_setup() 

45 self.commit() 

46 

47 # 3: function level changes 

48 

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)") 

56 

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))") 

63 

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))") 

67 

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() 

91 

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"] 

100 

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}):", ":") 

109 

110 # repl = f'class {match.expansions["$klass"][0]}:\n{raw(match.expansions["$$test_cases"])}' 

111 self.replace(repl, match.nodes, False, False) 

112 

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) 

120 

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) 

134 

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 

138 

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(") 

163 

164 self.replace(repl, fun, False, False) 

165 

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) 

174 

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) 

186 

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) 

193 

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) 

204 

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() 

233 

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}" 

243 

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) 

249 

250 for match in duplicate_imports[1:-1]: 

251 self.remove(match.nodes, False, False)