diff --git a/src/hio/core/tcp/clienting.py b/src/hio/core/tcp/clienting.py index 35378d8..2056190 100644 --- a/src/hio/core/tcp/clienting.py +++ b/src/hio/core/tcp/clienting.py @@ -252,12 +252,19 @@ def shutdown(self, how=socket.SHUT_RDWR): """ Shutdown connected socket .cs """ + shut = False # only record cutoff after successful shutdown if self.cs: try: self.cs.shutdown(how) # shutdown socket + shut = True except OSError as ex: pass + if shut and how in (socket.SHUT_RD, socket.SHUT_RDWR): + self.cutoff = True + if shut and how in (socket.SHUT_WR, socket.SHUT_RDWR): + self.txCutoff = True + def shutdownSend(self): """ @@ -281,6 +288,33 @@ def shutdownReceive(self): pass + def serviceClose(self): + """ + Service recurrent socket close. + Returns True when closed, False when service must retry. + """ + if not self.cs: + return True + + if self.txbs and not self.txCutoff: + return False # caller must settle output before write shutdown + + if not self.txCutoff: + try: + self.cs.shutdown(socket.SHUT_WR) + except OSError as ex: + self.txCutoff = True # failed shutdown is terminal for writes + self.error = ex + raise + self.txCutoff = True + + if not self.cutoff: + return False # wait for peer EOF after write shutdown + + self.close() + return self.cs is None + + def close(self): """ Shutdown and close connected socket .cs @@ -612,10 +646,10 @@ def connected(self, value): def close(self): """ - Shutdown and close connected socket .cs + Force close connected TLS socket .cs """ if self.cs: - self.shutdown() + # force close bypasses the recurrent TLS close_notify exchange self.cs.close() #close socket self.cs = None self.accepted = False diff --git a/src/hio/core/tcp/serving.py b/src/hio/core/tcp/serving.py index 190f4aa..ca55e88 100644 --- a/src/hio/core/tcp/serving.py +++ b/src/hio/core/tcp/serving.py @@ -725,6 +725,33 @@ def shutdownReceive(self): pass + def serviceClose(self): + """ + Service recurrent socket close. + Returns True when closed, False when service must retry. + """ + if not self.cs: + return True + + if self.txbs and not self.txCutoff: + return False # caller must settle output before write shutdown + + if not self.txCutoff: + try: + self.cs.shutdown(socket.SHUT_WR) + except OSError as ex: + self.txCutoff = True # failed shutdown is terminal for writes + self.error = ex + raise + self.txCutoff = True + + if not self.cutoff: + return False # wait for peer EOF after write shutdown + + self.close() + return self.cs is None + + def close(self): """ Shutdown and close connected socket .cs @@ -941,10 +968,10 @@ def __init__(self, def close(self): """ - Shutdown and close connected socket .cs + Force close connected TLS socket .cs """ if self.cs: - self.shutdown() + # force close bypasses the recurrent TLS close_notify exchange self.cs.close() #close socket self.cs = None self.connected = False diff --git a/tests/core/tcp/test_tcp.py b/tests/core/tcp/test_tcp.py index 1ab70f9..359277d 100644 --- a/tests/core/tcp/test_tcp.py +++ b/tests/core/tcp/test_tcp.py @@ -990,6 +990,19 @@ def test_tls_send_force_closes_connection_reset(endpointCls): assert endpoint.cs is None +@pytest.mark.parametrize("endpointCls", (tcp.ClientTls, serving.RemoterTls)) +def test_tls_force_close_does_not_raw_shutdown(endpointCls): + """Force-closing TLS bypasses raw socket shutdown.""" + cs = Mock(spec=ssl.SSLSocket) + endpoint = makeTlsEndpoint(endpointCls, cs) + + endpoint.close() + + assert endpoint.cs is None + cs.shutdown.assert_not_called() + cs.close.assert_called_once_with() + + @pytest.mark.parametrize("endpointCls", (tcp.ClientTls, serving.RemoterTls)) def test_tls_clean_close_ends_only_receive(endpointCls): """A peer close_notify cleanly ends only the receive direction.""" @@ -1039,6 +1052,80 @@ def test_tls_abrupt_eof_is_truncation(receiverName): closeTlsPair(server, client, remoter) +def test_client_shutdown_sets_directional_cutoffs(): + """Successful local shutdown records only the affected directions.""" + cs, peer = socket.socketpair() + client = tcp.Client(ha=("127.0.0.1", 6101)) + client.cs = cs + client.accepted = True + + try: + client.shutdownReceive() + assert client.cutoff is True + assert client.txCutoff is False + + client.shutdownSend() + assert client.cutoff is True + assert client.txCutoff is True + finally: + client.close() + peer.close() + + +@pytest.mark.parametrize("endpointCls", (tcp.Client, serving.Remoter)) +def test_tcp_service_close_waits_for_egress_and_peer_eof(endpointCls): + """Raw recurrent close drains output before awaiting peer EOF.""" + cs = Mock(spec=socket.socket) + if endpointCls is tcp.Client: + endpoint = tcp.Client(ha=("127.0.0.1", 6101)) + endpoint.cs = cs + endpoint.accepted = True + else: + endpoint = serving.Remoter(ha=("127.0.0.1", 6101), + ca=("127.0.0.1", 6102), + cs=cs) + + endpoint.tx(b"pending") + assert endpoint.serviceClose() is False + assert endpoint.txCutoff is False + cs.shutdown.assert_not_called() + + endpoint.txbs.clear() # owner reports egress settled + assert endpoint.serviceClose() is False + assert endpoint.txCutoff is True + cs.shutdown.assert_called_once_with(socket.SHUT_WR) + + endpoint.cutoff = True # peer EOF completes receive direction + assert endpoint.serviceClose() is True + assert endpoint.cs is None + cs.close.assert_called_once_with() + + +@pytest.mark.parametrize("endpointCls", (tcp.Client, serving.Remoter)) +def test_tcp_service_close_records_shutdown_error(endpointCls): + """Raw recurrent close exposes a terminal write-shutdown failure.""" + cs = Mock(spec=socket.socket) + failure = OSError(errno.ENOTCONN, "not connected") + cs.shutdown.side_effect = failure + if endpointCls is tcp.Client: + endpoint = tcp.Client(ha=("127.0.0.1", 6101)) + endpoint.cs = cs + endpoint.accepted = True + else: + endpoint = serving.Remoter(ha=("127.0.0.1", 6101), + ca=("127.0.0.1", 6102), + cs=cs) + + with pytest.raises(OSError) as excinfo: + endpoint.serviceClose() + + assert excinfo.value is failure + assert endpoint.error is failure + assert endpoint.cutoff is False + assert endpoint.txCutoff is True + assert endpoint.cs is cs + + def test_client_tracks_terminal_send_separately(): """Client broken pipe preserves receive service and unsent output.""" client = tcp.Client(ha=("127.0.0.1", 6101))