Skip to content

Commit 5692dfc

Browse files
authored
fix(websockets): Send close frame on ASGI return (#2769)
1 parent 4194764 commit 5692dfc

4 files changed

Lines changed: 27 additions & 13 deletions

File tree

tests/protocols/test_websocket.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -458,6 +458,27 @@ async def app(scope: Scope, receive: ASGIReceiveCallable, send: ASGISendCallable
458458
assert websocket.close_code == 1006
459459

460460

461+
async def test_close_transport_on_asgi_return(
462+
ws_protocol_cls: WSProtocol, http_protocol_cls: HTTPProtocol, unused_tcp_port: int
463+
):
464+
"""The ASGI callable should call the `websocket.close` event.
465+
466+
If it doesn't, the server should still send a close frame to the client.
467+
"""
468+
469+
async def app(scope: Scope, receive: ASGIReceiveCallable, send: ASGISendCallable):
470+
message = await receive()
471+
if message["type"] == "websocket.connect":
472+
await send({"type": "websocket.accept"})
473+
474+
config = Config(app=app, ws=ws_protocol_cls, http=http_protocol_cls, lifespan="off", port=unused_tcp_port)
475+
async with run_server(config):
476+
async with websockets.client.connect(f"ws://127.0.0.1:{unused_tcp_port}") as websocket:
477+
with pytest.raises(websockets.exceptions.ConnectionClosed):
478+
await websocket.recv()
479+
assert websocket.close_code == 1006
480+
481+
461482
@pytest.mark.parametrize("code", [None, 1000, 1001])
462483
@pytest.mark.parametrize("reason", [None, "test", False], ids=["none_as_reason", "normal_reason", "without_reason"])
463484
async def test_app_close(

uvicorn/protocols/websockets/websockets_impl.py

Lines changed: 2 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -244,25 +244,22 @@ async def run_asgi(self) -> None:
244244
result = await self.app(self.scope, self.asgi_receive, self.asgi_send) # type: ignore[func-returns-value]
245245
except ClientDisconnected: # pragma: full coverage
246246
self.closed_event.set()
247-
self.transport.close()
248247
except BaseException:
249248
self.closed_event.set()
250249
self.logger.exception("Exception in ASGI application\n")
251250
if not self.handshake_started_event.is_set():
252251
self.send_500_response()
253252
else:
254253
await self.handshake_completed_event.wait()
255-
self.transport.close()
256254
else:
257255
self.closed_event.set()
258256
if not self.handshake_started_event.is_set():
259257
self.logger.error("ASGI callable returned without sending handshake.")
260258
self.send_500_response()
261-
self.transport.close()
262259
elif result is not None:
263260
self.logger.error("ASGI callable should return None, but returned '%s'.", result)
264-
await self.handshake_completed_event.wait()
265-
self.transport.close()
261+
await self.handshake_completed_event.wait()
262+
self.transport.close()
266263

267264
async def asgi_send(self, message: ASGISendEvent) -> None:
268265
message_type = message["type"]

uvicorn/protocols/websockets/websockets_sansio_impl.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -271,19 +271,17 @@ async def run_asgi(self) -> None:
271271
try:
272272
result = await self.app(self.scope, self.receive, self.send)
273273
except ClientDisconnected:
274-
self.transport.close() # pragma: no cover
274+
pass # pragma: full coverage
275275
except BaseException:
276276
self.logger.exception("Exception in ASGI application\n")
277277
self.send_500_response()
278-
self.transport.close()
279278
else:
280279
if not self.handshake_complete:
281280
self.logger.error("ASGI callable returned without completing handshake.")
282281
self.send_500_response()
283-
self.transport.close()
284282
elif result is not None:
285283
self.logger.error("ASGI callable should return None, but returned '%s'.", result)
286-
self.transport.close()
284+
self.transport.close()
287285

288286
def send_500_response(self) -> None:
289287
if self.initial_response or self.handshake_complete:

uvicorn/protocols/websockets/wsproto_impl.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -234,19 +234,17 @@ async def run_asgi(self) -> None:
234234
try:
235235
result = await self.app(self.scope, self.receive, self.send) # type: ignore[func-returns-value]
236236
except ClientDisconnected:
237-
self.transport.close() # pragma: full coverage
237+
pass # pragma: full coverage
238238
except BaseException:
239239
self.logger.exception("Exception in ASGI application\n")
240240
self.send_500_response()
241-
self.transport.close()
242241
else:
243242
if not self.handshake_complete:
244243
self.logger.error("ASGI callable returned without completing handshake.")
245244
self.send_500_response()
246-
self.transport.close()
247245
elif result is not None:
248246
self.logger.error("ASGI callable should return None, but returned '%s'.", result)
249-
self.transport.close()
247+
self.transport.close()
250248

251249
async def send(self, message: ASGISendEvent) -> None:
252250
await self.writable.wait()

0 commit comments

Comments
 (0)