Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 36 additions & 2 deletions src/hio/core/tcp/clienting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
"""
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down
31 changes: 29 additions & 2 deletions src/hio/core/tcp/serving.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
87 changes: 87 additions & 0 deletions tests/core/tcp/test_tcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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))
Expand Down
Loading