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