gh-156920, gh-156698: fix ProactorEventLoop datagram transport hangs on close() and after write errors (#156921)

Co-authored-by: Kumar Aditya <kumaraditya@python.org>
diff --git a/Lib/asyncio/proactor_events.py b/Lib/asyncio/proactor_events.py
index 6717f06..8350e02 100644
--- a/Lib/asyncio/proactor_events.py
+++ b/Lib/asyncio/proactor_events.py
@@ -105,8 +105,9 @@ def close(self):
         if self._closing:
             return
         self._closing = True
-        self._conn_lost += 1
         if not self._buffer and self._write_fut is None:
+            # Nothing left to flush: no more data will be sent.
+            self._conn_lost += 1
             self._loop.call_soon(self._call_connection_lost, None)
         if self._read_fut is not None:
             self._read_fut.cancel()
@@ -386,6 +387,7 @@ def _loop_writing(self, f=None, data=None):
                 self._buffer = None
             if not data:
                 if self._closing:
+                    self._conn_lost += 1
                     self._loop.call_soon(self._call_connection_lost, None)
                 if self._eof_written:
                     self._sock.shutdown(socket.SHUT_WR)
@@ -480,6 +482,11 @@ def get_write_buffer_size(self):
     def abort(self):
         self._force_close(None)
 
+    def _force_close(self, exc):
+        # The base class drops the buffer; the size is tracked separately.
+        self._buffer_size = 0
+        super()._force_close(exc)
+
     def sendto(self, data, addr=None):
         if not isinstance(data, (bytes, bytearray, memoryview)):
             raise TypeError('data argument must be bytes-like object (%r)',
@@ -509,6 +516,8 @@ def sendto(self, data, addr=None):
     def _loop_writing(self, fut=None):
         try:
             if self._conn_lost:
+                # No more data will be sent: either everything buffered has
+                # already been flushed, or _force_close() dropped it.
                 return
 
             assert fut is self._write_fut
@@ -517,9 +526,10 @@ def _loop_writing(self, fut=None):
                 # We are in a _loop_writing() done callback, get the result
                 fut.result()
 
-            if not self._buffer or (self._conn_lost and self._address):
-                # The connection has been closed
+            if not self._buffer:
+                # Everything buffered has been sent
                 if self._closing:
+                    self._conn_lost += 1
                     self._loop.call_soon(self._call_connection_lost, None)
                 return
 
@@ -534,6 +544,27 @@ def _loop_writing(self, fut=None):
                                                               addr=addr)
         except OSError as exc:
             self._protocol.error_received(exc)
+            # error_received() is arbitrary protocol code: it may have sent
+            # (scheduling a write of its own, directly or via call_soon()),
+            # closed, or aborted the transport.
+            if self._buffer or self._closing:
+                # Either data is still queued, or a close() is waiting on
+                # the write loop to drain it and call connection_lost().
+                # This write failed, so there is no completion callback
+                # pending to re-enter the loop -- schedule one (gh-156698).
+                def write_next():
+                    # error_received() may have scheduled a write of its own,
+                    # directly or with call_soon(); its completion callback
+                    # will drain the rest of the buffer.
+                    if self._write_fut is None:
+                        self._loop_writing()
+
+                self._loop.call_soon(write_next)
+            else:
+                # Nothing left to write, so a paused protocol has to be
+                # resumed here: the next entry into _loop_writing() returns
+                # early on an empty buffer without doing it.
+                self._maybe_resume_protocol()
         except Exception as exc:
             self._fatal_error(exc, 'Fatal write error on datagram transport')
         else:
@@ -543,28 +574,20 @@ def _loop_writing(self, fut=None):
     def _loop_reading(self, fut=None):
         data = None
         try:
-            if self._conn_lost:
+            if self._closing:
                 return
 
-            assert self._read_fut is fut or (self._read_fut is None and
-                                             self._closing)
+            assert self._read_fut is fut
 
             self._read_fut = None
             if fut is not None:
                 res = fut.result()
 
-                if self._closing:
-                    # since close() has been called we ignore any read data
-                    data = None
-                    return
-
                 if self._address is not None:
                     data, addr = res, self._address
                 else:
                     data, addr = res
 
-            if self._conn_lost:
-                return
             if self._address is not None:
                 self._read_fut = self._loop._proactor.recv(self._sock,
                                                            self.max_size)
diff --git a/Lib/test/test_asyncio/test_events.py b/Lib/test/test_asyncio/test_events.py
index 06f538b..6368f4b 100644
--- a/Lib/test/test_asyncio/test_events.py
+++ b/Lib/test/test_asyncio/test_events.py
@@ -1583,6 +1583,264 @@ def create_socket():
         transport_1.close()
         transport_2.close()
 
+    def test_datagram_write_error_resumes_paused_protocol(self):
+        # See https://github.com/python/cpython/issues/156698: a
+        # datagram write error must not strand data left in the write
+        # buffer, nor leave a paused protocol paused forever.
+        loop = self.loop
+
+        class Protocol(asyncio.DatagramProtocol):
+            def connection_made(self, transport):
+                self.transport = transport
+                self.paused = False
+                self.resumed = False
+                self.errors = []
+                self.error_received_event = loop.create_future()
+
+            def pause_writing(self):
+                self.paused = True
+
+            def resume_writing(self):
+                self.resumed = True
+
+            def error_received(self, exc):
+                self.errors.append(exc)
+                if not self.error_received_event.done():
+                    self.error_received_event.set_result(None)
+
+        sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
+        sock.setblocking(False)
+        sock.bind(('127.0.0.1', 0))
+        transport, protocol = loop.run_until_complete(
+            loop.create_datagram_endpoint(Protocol, sock=sock))
+        addr = sock.getsockname()
+
+        # A high water mark of 0 makes pausing deterministic whenever
+        # anything is left in the write buffer.
+        transport.set_write_buffer_limits(0)
+
+        # The oversized datagram fails while it is in flight, and the
+        # normal datagram behind it is left queued -- queuing is also
+        # what trips pause_writing() at a high water mark of 0.
+        transport.sendto(b'\x00' * 70000, addr)
+        transport.sendto(b'queued', addr)
+
+        loop.run_until_complete(
+            asyncio.wait_for(protocol.error_received_event,
+                             support.SHORT_TIMEOUT))
+        self.assertTrue(protocol.errors)
+        self.assertIsInstance(protocol.errors[0], OSError)
+
+        # The write buffer must not be left stranded.
+        test_utils.run_until(
+            loop, lambda: transport.get_write_buffer_size() == 0)
+
+        # A protocol that got paused must eventually be resumed too --
+        # without requiring an unsolicited extra sendto() to un-stick it.
+        if protocol.paused:
+            test_utils.run_until(loop, lambda: protocol.resumed)
+
+        transport.close()
+        test_utils.run_briefly(loop)
+
+    def test_datagram_write_error_reentrant_sendto(self):
+        # See https://github.com/python/cpython/issues/156698: an
+        # error_received() callback that sends more data synchronously
+        # can itself schedule a new write. The write-loop restart scheduled
+        # for the failed write must notice that and not try to start a
+        # second, conflicting one.
+        loop = self.loop
+        unhandled = []
+        loop.set_exception_handler(lambda loop, context: unhandled.append(context))
+
+        class Protocol(asyncio.DatagramProtocol):
+            def connection_made(self, transport):
+                self.transport = transport
+                self.sent_extra = False
+                self.errors = []
+                self.done = loop.create_future()
+
+            def datagram_received(self, data, addr):
+                if not self.done.done():
+                    self.done.set_result(None)
+
+            def error_received(self, exc):
+                self.errors.append(exc)
+                if not self.sent_extra:
+                    # Reentrantly kicks off another write while the
+                    # failing one is still unwinding on the stack.
+                    self.sent_extra = True
+                    self.transport.sendto(b'extra', self.addr)
+
+        sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
+        sock.setblocking(False)
+        sock.bind(('127.0.0.1', 0))
+        transport, protocol = loop.run_until_complete(
+            loop.create_datagram_endpoint(Protocol, sock=sock))
+        protocol.addr = addr = sock.getsockname()
+
+        oversized = b'\x00' * 70000
+        transport.sendto(oversized, addr)
+        transport.sendto(b'queued', addr)
+
+        # The 'extra' datagram sent from error_received() is delivered
+        # back to the same socket; waiting for it proves the write loop
+        # kept running instead of wedging or crashing.
+        loop.run_until_complete(
+            asyncio.wait_for(protocol.done, support.SHORT_TIMEOUT))
+
+        test_utils.run_until(
+            loop, lambda: transport.get_write_buffer_size() == 0)
+
+        transport.close()
+        test_utils.run_briefly(loop)
+
+        self.assertTrue(protocol.errors)
+        self.assertFalse(
+            unhandled,
+            f'unhandled exception in the write loop: {unhandled}')
+
+    def test_datagram_close_flushes_queued_data(self):
+        # See https://github.com/python/cpython/issues/156920: _conn_lost
+        # used to mean "close() was requested" rather than "no more data
+        # will be sent". Since add_done_callback() always defers an
+        # already-completed write's callback with call_soon(), a sendto()
+        # immediately followed by close() -- with no await in between --
+        # leaves a write genuinely outstanding at close() time on every
+        # platform, not just a slow one. Closing must let that write (and
+        # anything queued behind it) drain and still call connection_lost(),
+        # instead of tripping the "no more data will be sent" guard before
+        # the drain has actually happened and hanging forever.
+        loop = self.loop
+
+        class Receiver(asyncio.DatagramProtocol):
+            def connection_made(self, transport):
+                self.received = []
+
+            def datagram_received(self, data, addr):
+                self.received.append(data)
+
+        recv_sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
+        recv_sock.setblocking(False)
+        recv_sock.bind(('127.0.0.1', 0))
+        recv_transport, receiver = loop.run_until_complete(
+            loop.create_datagram_endpoint(Receiver, sock=recv_sock))
+        addr = recv_sock.getsockname()
+
+        class Protocol(asyncio.DatagramProtocol):
+            def connection_made(self, transport):
+                self.lost = loop.create_future()
+
+            def connection_lost(self, exc):
+                if not self.lost.done():
+                    self.lost.set_result(exc)
+
+        sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
+        sock.setblocking(False)
+        sock.bind(('127.0.0.1', 0))
+        transport, protocol = loop.run_until_complete(
+            loop.create_datagram_endpoint(Protocol, sock=sock))
+
+        # 'first' is still in flight (its completion callback hasn't run
+        # yet) and 'second' is queued behind it when close() is called.
+        transport.sendto(b'first', addr)
+        transport.sendto(b'second', addr)
+        transport.close()
+
+        loop.run_until_complete(
+            asyncio.wait_for(protocol.lost, support.SHORT_TIMEOUT))
+
+        test_utils.run_until(
+            loop, lambda: len(receiver.received) >= 2)
+        self.assertEqual(sorted(receiver.received), [b'first', b'second'])
+
+        recv_transport.close()
+        test_utils.run_briefly(loop)
+
+    def test_datagram_close_during_write_error_calls_connection_lost(self):
+        # See https://github.com/python/cpython/issues/156920: if the
+        # write that's outstanding when close() is called goes on to fail
+        # (rather than succeed), the failure handler used to only re-schedule
+        # the write loop when data was still queued behind it. If that
+        # failing write was the last thing in the buffer, nothing re-scheduled
+        # the loop, so the close() in progress never got to call
+        # connection_lost() -- it hung forever instead of finishing once
+        # the buffer was actually empty.
+        loop = self.loop
+
+        class Protocol(asyncio.DatagramProtocol):
+            def connection_made(self, transport):
+                self.lost = loop.create_future()
+                self.errors = []
+
+            def error_received(self, exc):
+                self.errors.append(exc)
+
+            def connection_lost(self, exc):
+                if not self.lost.done():
+                    self.lost.set_result(exc)
+
+        sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
+        sock.setblocking(False)
+        sock.bind(('127.0.0.1', 0))
+        transport, protocol = loop.run_until_complete(
+            loop.create_datagram_endpoint(Protocol, sock=sock))
+        addr = sock.getsockname()
+
+        # 'ok' is still in flight when close() is called; 'oversized' is
+        # queued behind it and fails once it reaches the front of the
+        # buffer, leaving the buffer empty right as the error is handled.
+        oversized = b'\x00' * 70000
+        transport.sendto(b'ok', addr)
+        transport.sendto(oversized, addr)
+        transport.close()
+
+        loop.run_until_complete(
+            asyncio.wait_for(protocol.lost, support.SHORT_TIMEOUT))
+        self.assertTrue(protocol.errors)
+
+    def test_datagram_write_error_close_from_callback(self):
+        # See https://github.com/python/cpython/issues/156920: an
+        # error_received() callback that closes the transport must still
+        # result in connection_lost() being called eventually, instead of
+        # leaving the transport hanging forever. Two failing writes are
+        # used so that the first failure's error_received() call closes
+        # the transport while the second is still queued (close() defers
+        # to the write loop), and the second failure then empties the
+        # buffer with self._closing already True and no write in flight
+        # -- exercising the `self._closing` half of the
+        # `if self._buffer or self._closing:` condition in _loop_writing.
+        loop = self.loop
+
+        class Protocol(asyncio.DatagramProtocol):
+            def connection_made(self, transport):
+                self.transport = transport
+                self.errors = []
+                self.lost = loop.create_future()
+
+            def error_received(self, exc):
+                self.errors.append(exc)
+                self.transport.close()
+
+            def connection_lost(self, exc):
+                if not self.lost.done():
+                    self.lost.set_result(exc)
+
+        sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
+        sock.setblocking(False)
+        sock.bind(('127.0.0.1', 0))
+        transport, protocol = loop.run_until_complete(
+            loop.create_datagram_endpoint(Protocol, sock=sock))
+        addr = sock.getsockname()
+
+        oversized = b'\x00' * 70000
+        transport.sendto(oversized, addr)
+        transport.sendto(oversized, addr)
+
+        loop.run_until_complete(
+            asyncio.wait_for(protocol.lost, support.SHORT_TIMEOUT))
+        self.assertEqual(len(protocol.errors), 2)
+
     def test_datagram_recvfrom_connection_reset_recovers(self):
         # gh-127057: a UDP socket that sent a datagram to an address that
         # wasn't listening can raise ConnectionResetError on a later
diff --git a/Misc/NEWS.d/next/Library/2026-08-31-00-00-00.gh-issue-156698.gyDxUe.rst b/Misc/NEWS.d/next/Library/2026-08-31-00-00-00.gh-issue-156698.gyDxUe.rst
new file mode 100644
index 0000000..4e3a292
--- /dev/null
+++ b/Misc/NEWS.d/next/Library/2026-08-31-00-00-00.gh-issue-156698.gyDxUe.rst
@@ -0,0 +1,4 @@
+Fix :class:`asyncio.ProactorEventLoop` UDP transports so that a write
+error no longer strands a paused protocol: the write loop is now
+rescheduled when data remains buffered, and the protocol is resumed
+when the buffer has drained.
diff --git a/Misc/NEWS.d/next/Library/2026-09-04-09-18-05.gh-issue-156920.lONsKT.rst b/Misc/NEWS.d/next/Library/2026-09-04-09-18-05.gh-issue-156920.lONsKT.rst
new file mode 100644
index 0000000..f6bcf0c
--- /dev/null
+++ b/Misc/NEWS.d/next/Library/2026-09-04-09-18-05.gh-issue-156920.lONsKT.rst
@@ -0,0 +1,5 @@
+Fix :mod:`asyncio` on Windows: closing a :class:`~asyncio.DatagramTransport`
+under :class:`~asyncio.ProactorEventLoop` while datagrams were still queued,
+or while an in-flight write failed right as ``close()`` was draining the
+buffer, could strand the queued data and never call ``connection_lost()``,
+hanging the close indefinitely.