diff --git a/Lib/test/test_codeop.py b/Lib/test/test_codeop.py index ed10bd3dcb6d2b..d57452602ce557 100644 --- a/Lib/test/test_codeop.py +++ b/Lib/test/test_codeop.py @@ -4,43 +4,64 @@ """ import unittest import warnings -from test.support import warnings_helper +from test.support import subTests, warnings_helper from textwrap import dedent +import functools -from codeop import compile_command, PyCF_DONT_IMPLY_DEDENT +from codeop import compile_command, CommandCompiler, Compile +from codeop import PyCF_DONT_IMPLY_DEDENT, PyCF_ONLY_AST +import ast + + +WRAPPING_COMPILERS = [compile_command, CommandCompiler()] +RAW_COMPILERS = [Compile()] +COMPILERS = WRAPPING_COMPILERS + RAW_COMPILERS -class CodeopTests(unittest.TestCase): - def assertValid(self, str, symbol='single'): +class CodeopTests(unittest.TestCase): + def assertValid(self, str, symbol='single', *, compiler): '''succeed iff str is a valid piece of code''' expected = compile(str, "", symbol, PyCF_DONT_IMPLY_DEDENT) - self.assertEqual(compile_command(str, "", symbol), expected) + self.assertEqual(compiler(str, "", symbol), expected) - def assertIncomplete(self, str, symbol='single'): + def assertIncomplete(self, str, symbol='single', *, compiler): '''succeed iff str is the start of a valid piece of code''' - self.assertEqual(compile_command(str, symbol=symbol), None) - - def assertInvalid(self, str, symbol='single', is_syntax=1): + if compiler in WRAPPING_COMPILERS: + self.assertEqual(compiler(str, "", symbol=symbol), None) + else: + # Compile has should raise like built-in compile + with self.assertRaises(SyntaxError) as cm_original_error: + compile(str, "", symbol, compiler.flags) + expected_error = cm_original_error.exception + with self.assertRaises(type(expected_error)) as cm_wrapped_error: + compiler(str, "", symbol=symbol) + self.assertEqual( + expected_error.args, + cm_wrapped_error.exception.args + ) + + def assertInvalid(self, str, symbol='single', is_syntax=1, *, compiler): '''succeed iff str is the start of an invalid piece of code''' try: - compile_command(str,symbol=symbol) + compiler(str,"", symbol=symbol) self.fail("No exception raised for invalid code") except SyntaxError: self.assertTrue(is_syntax) except OverflowError: self.assertTrue(not is_syntax) - def test_valid(self): - av = self.assertValid - - # special case - self.assertEqual(compile_command(""), - compile("pass", "", 'single', - PyCF_DONT_IMPLY_DEDENT)) - self.assertEqual(compile_command("\n"), - compile("pass", "", 'single', - PyCF_DONT_IMPLY_DEDENT)) - + @subTests('compiler', WRAPPING_COMPILERS) + def test_empty(self, compiler): + self.assertEqual( + compiler("", "", 'single'), + compile("pass", "", 'single', PyCF_DONT_IMPLY_DEDENT)) + self.assertEqual( + compiler("\n", "", 'single'), + compile("pass", "", 'single', PyCF_DONT_IMPLY_DEDENT)) + + @subTests('compiler', COMPILERS) + def test_valid(self, compiler): + av = functools.partial(self.assertValid, compiler=compiler) av("a = 1") av("\na = 1") av("a = 1\n") @@ -92,8 +113,9 @@ def test_valid(self): av("def f():\n pass\n#foo\n") av("@a.b.c\ndef f():\n pass\n") - def test_incomplete(self): - ai = self.assertIncomplete + @subTests('compiler', COMPILERS) + def test_incomplete(self, compiler): + ai = functools.partial(self.assertIncomplete, compiler=compiler) ai("(a **") ai("(a,b,") @@ -226,8 +248,9 @@ def test_incomplete(self): ai('a = f"""') ai('a = \\') - def test_invalid(self): - ai = self.assertInvalid + @subTests('compiler', COMPILERS) + def test_invalid(self, compiler): + ai = functools.partial(self.assertInvalid, compiler=compiler) ai("a b") ai("a @") @@ -263,8 +286,9 @@ def test_invalid(self): ai("[i for i in range(10)] = (1, 2, 3)") - def test_invalid_exec(self): - ai = self.assertInvalid + @subTests('compiler', COMPILERS) + def test_invalid_exec(self, compiler): + ai = functools.partial(self.assertInvalid, compiler=compiler) ai("raise = 4", symbol="exec") ai('def a-b', symbol='exec') ai('await?', symbol='exec') @@ -272,58 +296,94 @@ def test_invalid_exec(self): ai('a await raise b', symbol='exec') ai('a await raise b?+1', symbol='exec') - def test_filename(self): - self.assertEqual(compile_command("a = 1\n", "abc").co_filename, - compile("a = 1\n", "abc", 'single').co_filename) - self.assertNotEqual(compile_command("a = 1\n", "abc").co_filename, - compile("a = 1\n", "def", 'single').co_filename) - - def test_warning(self): + @subTests('compiler', COMPILERS) + def test_filename(self, compiler): + self.assertEqual( + compiler("a = 1\n", "abc", "single").co_filename, + compile("a = 1\n", "abc", 'single').co_filename + ) + self.assertNotEqual( + compiler("a = 1\n", "abc", "single").co_filename, + compile("a = 1\n", "def", 'single').co_filename + ) + + def assertReturnsModule(self, code, compiler): + retval = compiler(code, "", 'exec', PyCF_ONLY_AST) + self.assertIsInstance(retval, ast.Module) + + @subTests('compiler', RAW_COMPILERS) + def test_ast_return_value(self, compiler): + validate_ast = self.assertReturnsModule + validate_ast("x = 5", compiler) + validate_ast("\nx = 5", compiler) + validate_ast("x = 5\n", compiler) + validate_ast("x = 5\n\n", compiler) + validate_ast("\n\nx = 5\n\n", compiler) + + @subTests('compiler', COMPILERS) + def test_warning(self, compiler): # Test that the warning is only returned once. with warnings_helper.check_warnings( ('"is" with \'str\' literal', SyntaxWarning), ('"\\\\e" is an invalid escape sequence', SyntaxWarning), ) as w: - compile_command(r"'\e' is 0") - self.assertEqual(len(w.warnings), 2) + compiler(r"'\e' is 0", "", "single") + self.assertEqual(len(w.warnings), 2) # bpo-41520: check SyntaxWarning treated as an SyntaxError with warnings.catch_warnings(), self.assertRaises(SyntaxError): warnings.simplefilter('error', SyntaxWarning) - compile_command('1 is 1', symbol='exec') + compiler('1 is 1', "", 'exec') # Check SyntaxWarning treated as an SyntaxError with warnings.catch_warnings(), self.assertRaises(SyntaxError): warnings.simplefilter('error', SyntaxWarning) - compile_command(r"'\e'", symbol='exec') + compiler(r"'\e'", "", 'exec') - def test_incomplete_warning(self): + @subTests('compiler', WRAPPING_COMPILERS) + def test_incomplete_warning(self, compiler): with warnings.catch_warnings(record=True) as w: warnings.simplefilter('always') - self.assertIncomplete("'\\e' + (") + compiler("'\\e' + (") self.assertEqual(w, []) - def test_invalid_warning(self): + @subTests('compiler', RAW_COMPILERS) + def test_raw_raises_error(self, compiler): + warnings_cm = warnings_helper.check_warnings( + ('"\\\\e" is an invalid esceape sequence', SyntaxWarning) + ) + with self.assertRaises(SyntaxError), warnings_cm as w: + compiler("'\\e' + (", "", 'single') + self.assertEqual(len(w.warnings), 1) + + @subTests('compiler', COMPILERS) + def test_invalid_warning(self, compiler): with warnings.catch_warnings(record=True) as w: warnings.simplefilter('always') - self.assertInvalid("'\\e' 1") + self.assertInvalid("'\\e' 1", compiler=compiler) self.assertEqual(len(w), 1) - self.assertEqual(w[0].category, SyntaxWarning) - self.assertRegex(str(w[0].message), 'invalid escape sequence') - self.assertEqual(w[0].filename, '') - - def assertSyntaxErrorMatches(self, code, message): - with self.subTest(code): - with self.assertRaisesRegex(SyntaxError, message): - compile_command(code, symbol='exec') - - def test_syntax_errors(self): - self.assertSyntaxErrorMatches( - dedent("""\ + for warning in w: + self.assertEqual(warning.category, SyntaxWarning) + self.assertRegex(str(warning), 'invalid escape sequence') + self.assertEqual(warning.filename, '') + + @subTests('compiler', COMPILERS) + def test_syntax_errors(self, compiler): + code = dedent("""\ def foo(x,x): pass - """), "duplicate parameter 'x' in function definition") - + """) + message = "duplicate parameter 'x' in function definition" + with self.assertRaisesRegex(SyntaxError, message): + compiler(code, "", 'exec') + + @subTests('compiler', RAW_COMPILERS) + def test_future_imports(self, compiler): + original_flags = compiler.flags + compiler('from __future__ import annotations', "", 'single') + self.assertGreater(compiler.flags, original_flags) + # reset flags to ensure test has no side-effects + compiler.flags = original_flags if __name__ == "__main__":