Skip to content

Commit e6f7ad1

Browse files
graingertKludex
andauthored
avoid collapsing exception groups from user code (#2830)
Co-authored-by: Marcelo Trylesinski <marcelotryle@gmail.com>
1 parent 115228f commit e6f7ad1

7 files changed

Lines changed: 84 additions & 33 deletions

File tree

starlette/_utils.py

Lines changed: 20 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -2,10 +2,12 @@
22

33
import functools
44
import sys
5-
from collections.abc import Awaitable, Callable, Generator
6-
from contextlib import AbstractAsyncContextManager, contextmanager
5+
from collections.abc import AsyncGenerator, Awaitable, Callable, Generator
6+
from contextlib import AbstractAsyncContextManager, asynccontextmanager
77
from typing import Any, Generic, Protocol, TypeVar, overload
88

9+
import anyio.abc
10+
911
from starlette.types import Scope
1012

1113
if sys.version_info >= (3, 13): # pragma: no cover
@@ -16,12 +18,14 @@
1618

1719
from typing_extensions import TypeIs
1820

19-
has_exceptiongroups = True
2021
if sys.version_info < (3, 11): # pragma: no cover
2122
try:
22-
from exceptiongroup import BaseExceptionGroup # type: ignore[unused-ignore,import-not-found]
23+
from exceptiongroup import BaseExceptionGroup
2324
except ImportError:
24-
has_exceptiongroups = False
25+
26+
class BaseExceptionGroup(BaseException): # type: ignore[no-redef]
27+
pass
28+
2529

2630
T = TypeVar("T")
2731
AwaitableCallable = Callable[..., Awaitable[T]]
@@ -75,16 +79,18 @@ async def __aexit__(self, *args: Any) -> None | bool:
7579
return None
7680

7781

78-
@contextmanager
79-
def collapse_excgroups() -> Generator[None, None, None]:
82+
@asynccontextmanager
83+
async def create_collapsing_task_group() -> AsyncGenerator[anyio.abc.TaskGroup, None]:
8084
try:
81-
yield
82-
except BaseException as exc:
83-
if has_exceptiongroups: # pragma: no cover
84-
while isinstance(exc, BaseExceptionGroup) and len(exc.exceptions) == 1:
85-
exc = exc.exceptions[0]
86-
87-
raise exc
85+
async with anyio.create_task_group() as tg:
86+
yield tg
87+
except BaseExceptionGroup as excs:
88+
if len(excs.exceptions) != 1:
89+
raise
90+
91+
exc = excs.exceptions[0]
92+
context = None if exc.__suppress_context__ else exc.__context__
93+
raise exc from exc.__cause__ or context
8894

8995

9096
def get_route_path(scope: Scope) -> str:

starlette/middleware/base.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55

66
import anyio
77

8-
from starlette._utils import collapse_excgroups
8+
from starlette._utils import create_collapsing_task_group
99
from starlette.requests import ClientDisconnect, Request
1010
from starlette.responses import Response
1111
from starlette.types import ASGIApp, Message, Receive, Scope, Send
@@ -188,8 +188,8 @@ async def body_stream() -> BodyStreamGenerator:
188188

189189
streams: anyio.create_memory_object_stream[Message] = anyio.create_memory_object_stream()
190190
send_stream, recv_stream = streams
191-
with recv_stream, send_stream, collapse_excgroups():
192-
async with anyio.create_task_group() as task_group:
191+
with recv_stream, send_stream:
192+
async with create_collapsing_task_group() as task_group:
193193
response = await self.dispatch_func(request, call_next)
194194
await response(scope, wrapped_receive, send)
195195
response_sent.set()

starlette/middleware/wsgi.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
import anyio
1111
from anyio.abc import ObjectReceiveStream, ObjectSendStream
1212

13+
from starlette._utils import create_collapsing_task_group
1314
from starlette.types import Receive, Scope, Send
1415

1516
warnings.warn(
@@ -104,7 +105,7 @@ async def __call__(self, receive: Receive, send: Send) -> None:
104105
more_body = message.get("more_body", False)
105106
environ = build_environ(self.scope, body)
106107

107-
async with anyio.create_task_group() as task_group:
108+
async with create_collapsing_task_group() as task_group:
108109
task_group.start_soon(self.sender, send)
109110
async with self.stream_send:
110111
await anyio.to_thread.run_sync(self.wsgi, environ, self.start_response)

starlette/responses.py

Lines changed: 7 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
import anyio
1919
import anyio.to_thread
2020

21-
from starlette._utils import collapse_excgroups
21+
from starlette._utils import create_collapsing_task_group
2222
from starlette.background import BackgroundTask
2323
from starlette.concurrency import iterate_in_threadpool
2424
from starlette.datastructures import URL, Headers, MutableHeaders
@@ -270,15 +270,14 @@ async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
270270
except OSError:
271271
raise ClientDisconnect()
272272
else:
273-
with collapse_excgroups():
274-
async with anyio.create_task_group() as task_group:
273+
async with create_collapsing_task_group() as task_group:
275274

276-
async def wrap(func: Callable[[], Awaitable[None]]) -> None:
277-
await func()
278-
task_group.cancel_scope.cancel()
275+
async def wrap(func: Callable[[], Awaitable[None]]) -> None:
276+
await func()
277+
task_group.cancel_scope.cancel()
279278

280-
task_group.start_soon(wrap, partial(self.stream_response, send))
281-
await wrap(partial(self.listen_for_disconnect, receive))
279+
task_group.start_soon(wrap, partial(self.stream_response, send))
280+
await wrap(partial(self.listen_for_disconnect, receive))
282281

283282
if self.background is not None:
284283
await self.background()

tests/middleware/test_base.py

Lines changed: 18 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from __future__ import annotations
22

33
import contextvars
4+
import sys
45
from collections.abc import AsyncGenerator, AsyncIterator, Generator
56
from contextlib import AsyncExitStack
67
from pathlib import Path
@@ -22,6 +23,9 @@
2223
from starlette.websockets import WebSocket
2324
from tests.types import TestClientFactory
2425

26+
if sys.version_info < (3, 11): # pragma: no cover
27+
from exceptiongroup import ExceptionGroup
28+
2529

2630
class CustomMiddleware(BaseHTTPMiddleware):
2731
async def dispatch(
@@ -42,6 +46,10 @@ def exc(request: Request) -> None:
4246
raise Exception("Exc")
4347

4448

49+
def exc_group(request: Request) -> None:
50+
raise ExceptionGroup("my exception group", [ValueError("TEST")])
51+
52+
4553
def exc_stream(request: Request) -> StreamingResponse:
4654
return StreamingResponse(_generate_faulty_stream())
4755

@@ -77,6 +85,7 @@ async def websocket_endpoint(session: WebSocket) -> None:
7785
routes=[
7886
Route("/", endpoint=homepage),
7987
Route("/exc", endpoint=exc),
88+
Route("/exc-group", endpoint=exc_group),
8089
Route("/exc-stream", endpoint=exc_stream),
8190
Route("/no-response", endpoint=NoResponse),
8291
WebSocketRoute("/ws", endpoint=websocket_endpoint),
@@ -90,13 +99,18 @@ def test_custom_middleware(test_client_factory: TestClientFactory) -> None:
9099
response = client.get("/")
91100
assert response.headers["Custom-Header"] == "Example"
92101

93-
with pytest.raises(Exception) as ctx:
102+
with pytest.raises(Exception) as ctx1:
94103
response = client.get("/exc")
95-
assert str(ctx.value) == "Exc"
104+
assert str(ctx1.value) == "Exc"
96105

97-
with pytest.raises(Exception) as ctx:
106+
with pytest.raises(Exception) as ctx2:
98107
response = client.get("/exc-stream")
99-
assert str(ctx.value) == "Faulty Stream"
108+
assert str(ctx2.value) == "Faulty Stream"
109+
110+
with pytest.raises(ExceptionGroup, match="my exception group") as ctx3:
111+
client.get("/exc-group")
112+
assert len(ctx3.value.exceptions) == 1
113+
assert isinstance(ctx3.value.exceptions[0], ValueError)
100114

101115
with pytest.raises(RuntimeError):
102116
response = client.get("/no-response")

tests/middleware/test_wsgi.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,6 @@
44

55
import pytest
66

7-
from starlette._utils import collapse_excgroups
87
from starlette.middleware.wsgi import WSGIMiddleware, build_environ
98
from tests.types import TestClientFactory
109

@@ -86,7 +85,7 @@ def test_wsgi_exception(test_client_factory: TestClientFactory) -> None:
8685
# The HTTP protocol implementations would catch this error and return 500.
8786
app = WSGIMiddleware(raise_exception)
8887
client = test_client_factory(app)
89-
with pytest.raises(RuntimeError), collapse_excgroups():
88+
with pytest.raises(RuntimeError):
9089
client.get("/")
9190

9291

tests/test__utils.py

Lines changed: 33 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,16 @@
11
import functools
2+
import sys
23
from typing import Any
34
from unittest.mock import create_autospec
45

56
import pytest
67

7-
from starlette._utils import get_route_path, is_async_callable
8+
from starlette._utils import create_collapsing_task_group, get_route_path, is_async_callable
89
from starlette.types import Scope
910

11+
if sys.version_info < (3, 11): # pragma: no cover
12+
from exceptiongroup import ExceptionGroup
13+
1014

1115
def test_async_func() -> None:
1216
async def async_func() -> None: ... # pragma: no cover
@@ -102,3 +106,31 @@ async def async_func() -> None: ... # pragma: no cover
102106
)
103107
def test_get_route_path(scope: Scope, expected_result: str) -> None:
104108
assert get_route_path(scope) == expected_result
109+
110+
111+
@pytest.mark.anyio
112+
async def test_collapsing_task_group_one_exc() -> None:
113+
class MyException(Exception):
114+
pass
115+
116+
with pytest.raises(MyException):
117+
async with create_collapsing_task_group():
118+
raise MyException
119+
120+
121+
@pytest.mark.anyio
122+
async def test_collapsing_task_group_two_exc() -> None:
123+
class MyException(Exception):
124+
pass
125+
126+
async def raise_exc() -> None:
127+
raise MyException
128+
129+
with pytest.raises(ExceptionGroup) as exc:
130+
async with create_collapsing_task_group() as task_group:
131+
task_group.start_soon(raise_exc)
132+
raise MyException
133+
134+
exc1, exc2 = exc.value.exceptions
135+
assert isinstance(exc1, MyException)
136+
assert isinstance(exc2, MyException)

0 commit comments

Comments
 (0)