diff --git a/docs/changelog.rst b/docs/changelog.rst index 8cf036e..7e43094 100644 --- a/docs/changelog.rst +++ b/docs/changelog.rst @@ -4,6 +4,10 @@ Changelog `CalVer, YY.month.patch `_ +Future +====== +- Extend :ref:`ASYNC401 ` to catch more exception-group assertions. `(issue #475) `_ + 26.8.1 ====== - Add :ref:`ASYNC128 ` task-status-never-started, warning about startable functions (i.e. with a ``task_status`` parameter) that never call ``task_status.started()``. `(issue #471) `_ diff --git a/docs/rules.rst b/docs/rules.rst index 2082bcc..2befc33 100644 --- a/docs/rules.rst +++ b/docs/rules.rst @@ -233,7 +233,7 @@ _`ASYNC400` : except-star-invalid-attribute When converting a codebase to use `except* ` it's easy to miss that the caught exception(s) are wrapped in a group, so accessing attributes on the caught exception must now check the contained exceptions. This checks for any attribute access on a caught ``except*`` that's not a known valid attribute on `ExceptionGroup`. This can be safely disabled on a type-checked or coverage-covered code base. _`ASYNC401` : pytest-raises-exception-group - ``pytest.raises(ExceptionGroup)`` and ``pytest.raises(BaseExceptionGroup)`` usually hide the structure of exception groups. Prefer ``pytest.RaisesGroup``. + ``pytest.raises(ExceptionGroup)`` and ``pytest.raises(BaseExceptionGroup)`` usually hide the structure of exception groups. Prefer ``pytest.RaisesGroup``. Similar pytest assertions are checked too. Optional rules disabled by default ================================== diff --git a/flake8_async/visitors/visitor4xx.py b/flake8_async/visitors/visitor4xx.py index 90cdf56..b4e3124 100644 --- a/flake8_async/visitors/visitor4xx.py +++ b/flake8_async/visitors/visitor4xx.py @@ -112,10 +112,7 @@ def visit_FunctionDef( @error_class class Visitor401(Flake8AsyncVisitor): error_codes: Mapping[str, str] = { - "ASYNC401": ( - "Use `pytest.RaisesGroup` instead of `pytest.raises({})` when expecting" - " exception groups." - ) + "ASYNC401": "Use `pytest.RaisesGroup` instead of expecting {} directly." } def _exception_group_name(self, node: ast.expr) -> str | None: @@ -125,7 +122,9 @@ def _exception_group_name(self, node: ast.expr) -> str | None: return name return None - canonical = self.canonical_name(node) + canonical = self.canonical_name( + node.value if isinstance(node, ast.Subscript) else node + ) if canonical in EXCGROUP_QUALNAMES: return ast.unparse(node) return None @@ -139,9 +138,22 @@ def _expected_exception_arg(self, node: ast.Call) -> ast.expr | None: return None def visit_Call(self, node: ast.Call): - if ( - self.canonical_name(node.func) == "pytest.raises" - and (expected_exception := self._expected_exception_arg(node)) is not None - and (exception_group := self._exception_group_name(expected_exception)) - ): - self.error(node, exception_group) + name = self.canonical_name(node.func) + if name == "pytest.mark.xfail": + expected_exceptions = [ + kw.value for kw in node.keywords if kw.arg == "raises" + ] + elif name == "pytest.RaisesGroup": + expected_exceptions = node.args + elif name in ("pytest.raises", "pytest.RaisesExc"): + expected_exception = self._expected_exception_arg(node) + expected_exceptions = ( + [] if expected_exception is None else [expected_exception] + ) + else: + return + + for expected_exception in expected_exceptions: + if exception_group := self._exception_group_name(expected_exception): + self.error(node, exception_group) + break diff --git a/tests/eval_files/async401.py b/tests/eval_files/async401.py index 0abaf85..ae028a0 100644 --- a/tests/eval_files/async401.py +++ b/tests/eval_files/async401.py @@ -4,6 +4,9 @@ import pytest from exceptiongroup import BaseExceptionGroup as BackportBaseExceptionGroup from exceptiongroup import ExceptionGroup as BackportExceptionGroup +from pytest import RaisesExc as raises_exc +from pytest import RaisesGroup as raises_group +from pytest import mark as pytest_mark from pytest import raises from pytest import raises as pytest_raises @@ -32,3 +35,42 @@ def raises(self, expected_exception): pytest.RaisesGroup(ValueError) raises(ValueError) not_pytest.raises(ExceptionGroup) + +pytest.raises(ExceptionGroup[Exception]) # error: 0, "ExceptionGroup[Exception]" +pytest.RaisesExc(ExceptionGroup) # error: 0, "ExceptionGroup" +pytest.RaisesExc( # error: 0, "BaseExceptionGroup" + expected_exception=BaseExceptionGroup +) +pytest.RaisesGroup(ExceptionGroup) # error: 0, "ExceptionGroup" +pytest.RaisesGroup(ValueError, BaseExceptionGroup) # error: 0, "BaseExceptionGroup" +pytest.mark.xfail(raises=ExceptionGroup) # error: 0, "ExceptionGroup" +pytest.mark.xfail( # error: 0, "ExceptionGroup" + reason="expected", raises=(ValueError, ExceptionGroup) +) + +pytest.RaisesExc(ValueError) +pytest.RaisesGroup(pytest.RaisesGroup(ValueError)) +pytest.mark.xfail(raises=ValueError) +pytest.mark.xfail(reason="expected") + +raises_exc(BackportExceptionGroup) # error: 0, "BackportExceptionGroup" +raises_group( # error: 0, "ExceptionGroup[Exception]" + ValueError, ExceptionGroup[Exception] +) +pytest_mark.xfail( # error: 0, "BackportBaseExceptionGroup" + raises=BackportBaseExceptionGroup +) +pytest.RaisesGroup(ExceptionGroup, BaseExceptionGroup) # error: 0, "ExceptionGroup" +pytest.RaisesGroup(pytest.RaisesGroup(ExceptionGroup)) # error: 19, "ExceptionGroup" +pytest.RaisesExc(builtins.BaseExceptionGroup) # type: ignore[attr-defined] # error: 0, "builtins.BaseExceptionGroup" + +pytest.raises((ValueError, TypeError)) +pytest.raises(match="message") +pytest.RaisesExc(match="message") +pytest.RaisesGroup() +pytest.RaisesGroup(ValueError, TypeError, match="message") +pytest.mark.xfail(ExceptionGroup) +pytest.mark.xfail(raises=pytest.RaisesGroup(ValueError)) +pytest.mark.xfail(**{"raises": ExceptionGroup}) +not_pytest.raises(ExceptionGroup[Exception]) +pytest.raises(list[ExceptionGroup])