Skip to content

Commit 633b4ed

Browse files
committed
gh-155305: Preserve user-provided sockets on transport errors
1 parent 7c653e2 commit 633b4ed

5 files changed

Lines changed: 43 additions & 10 deletions

File tree

Lib/asyncio/base_events.py

Lines changed: 11 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1083,6 +1083,7 @@ async def create_connection(
10831083
connection in the background. When successful, the coroutine
10841084
returns a (transport, protocol) pair.
10851085
"""
1086+
sock_was_provided = sock is not None
10861087
if server_hostname is not None and not ssl:
10871088
raise ValueError('server_hostname is only meaningful with ssl')
10881089

@@ -1204,7 +1205,8 @@ async def create_connection(
12041205
transport, protocol = await self._create_connection_transport(
12051206
sock, protocol_factory, ssl, server_hostname,
12061207
ssl_handshake_timeout=ssl_handshake_timeout,
1207-
ssl_shutdown_timeout=ssl_shutdown_timeout)
1208+
ssl_shutdown_timeout=ssl_shutdown_timeout,
1209+
sock_was_provided=sock_was_provided)
12081210
if self._debug:
12091211
# Get the socket from the transport because SSL transport closes
12101212
# the old socket and creates a new SSL socket
@@ -1217,7 +1219,8 @@ async def _create_connection_transport(
12171219
self, sock, protocol_factory, ssl,
12181220
server_hostname, server_side=False,
12191221
ssl_handshake_timeout=None,
1220-
ssl_shutdown_timeout=None, context=None):
1222+
ssl_shutdown_timeout=None, context=None,
1223+
sock_was_provided=False):
12211224

12221225
try:
12231226
sock.setblocking(False)
@@ -1236,8 +1239,10 @@ async def _create_connection_transport(
12361239
else:
12371240
transport = self._make_socket_transport(sock, protocol, waiter, context=context)
12381241
except:
1239-
# gh-153133: close the socket if the transport is never created.
1240-
sock.close()
1242+
# gh-153133: close internally created sockets if the transport is
1243+
# never created.
1244+
if not sock_was_provided:
1245+
sock.close()
12411246
raise
12421247

12431248
try:
@@ -1705,7 +1710,8 @@ async def connect_accepted_socket(
17051710
transport, protocol = await self._create_connection_transport(
17061711
sock, protocol_factory, ssl, '', server_side=True,
17071712
ssl_handshake_timeout=ssl_handshake_timeout,
1708-
ssl_shutdown_timeout=ssl_shutdown_timeout)
1713+
ssl_shutdown_timeout=ssl_shutdown_timeout,
1714+
sock_was_provided=True)
17091715
if self._debug:
17101716
# Get the socket from the transport because SSL transport closes
17111717
# the old socket and creates a new SSL socket

Lib/asyncio/unix_events.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -227,6 +227,7 @@ async def create_unix_connection(
227227
server_hostname=None,
228228
ssl_handshake_timeout=None,
229229
ssl_shutdown_timeout=None):
230+
sock_was_provided = sock is not None
230231
assert server_hostname is None or isinstance(server_hostname, str)
231232
if ssl:
232233
if server_hostname is None:
@@ -268,7 +269,8 @@ async def create_unix_connection(
268269
transport, protocol = await self._create_connection_transport(
269270
sock, protocol_factory, ssl, server_hostname,
270271
ssl_handshake_timeout=ssl_handshake_timeout,
271-
ssl_shutdown_timeout=ssl_shutdown_timeout)
272+
ssl_shutdown_timeout=ssl_shutdown_timeout,
273+
sock_was_provided=sock_was_provided)
272274
return transport, protocol
273275

274276
async def create_unix_server(

Lib/test/test_asyncio/test_base_events.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1341,9 +1341,9 @@ def getaddrinfo(*args, **kw):
13411341
self.loop.run_until_complete(coro)
13421342
self.assertTrue(sock.close.called)
13431343

1344-
def test_create_connection_sock_transport_error_closes_sock(self):
1345-
# gh-153133: a user-provided socket is closed if the transport is
1346-
# never created.
1344+
def test_create_connection_sock_transport_error_does_not_close_sock(self):
1345+
# gh-155305: a user-provided socket remains owned by the caller when
1346+
# the transport is never created.
13471347
sock = mock.Mock()
13481348
sock.type = socket.SOCK_STREAM
13491349

@@ -1353,7 +1353,7 @@ def factory():
13531353
coro = self.loop.create_connection(factory, sock=sock)
13541354
with self.assertRaises(ZeroDivisionError):
13551355
self.loop.run_until_complete(coro)
1356-
self.assertTrue(sock.close.called)
1356+
self.assertFalse(sock.close.called)
13571357

13581358
@patch_socket
13591359
def test_create_connection_transport_error_closes_sock(self, m_socket):

Lib/test/test_asyncio/test_events.py

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -914,6 +914,18 @@ def test_connect_accepted_socket_ssl_timeout_for_plain_socket(self):
914914
'ssl_handshake_timeout is only meaningful with ssl'):
915915
self.loop.run_until_complete(coro)
916916

917+
def test_connect_accepted_socket_transport_error_does_not_close_sock(self):
918+
sock = mock.Mock()
919+
sock.type = socket.SOCK_STREAM
920+
921+
def factory():
922+
raise ZeroDivisionError
923+
924+
coro = self.loop.connect_accepted_socket(factory, sock)
925+
with self.assertRaises(ZeroDivisionError):
926+
self.loop.run_until_complete(coro)
927+
self.assertFalse(sock.close.called)
928+
917929
@mock.patch('asyncio.base_events.socket')
918930
def create_server_multiple_hosts(self, family, hosts, mock_sock):
919931
async def getaddrinfo(host, port, *args, **kw):

Lib/test/test_asyncio/test_unix_events.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -393,6 +393,19 @@ def test_create_unix_connection_path_inetsock(self):
393393
'A UNIX Domain Stream.*was expected'):
394394
self.loop.run_until_complete(coro)
395395

396+
def test_create_unix_connection_transport_error_does_not_close_sock(self):
397+
sock = mock.Mock()
398+
sock.family = socket.AF_UNIX
399+
sock.type = socket.SOCK_STREAM
400+
401+
def factory():
402+
raise ZeroDivisionError
403+
404+
coro = self.loop.create_unix_connection(factory, sock=sock)
405+
with self.assertRaises(ZeroDivisionError):
406+
self.loop.run_until_complete(coro)
407+
self.assertFalse(sock.close.called)
408+
396409
@mock.patch('asyncio.unix_events.socket')
397410
def test_create_unix_server_bind_error(self, m_socket):
398411
# Ensure that the socket is closed on any bind error

0 commit comments

Comments
 (0)