Skip to content
Open
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
191 changes: 113 additions & 78 deletions ankaios_sdk/_components/connection/command_interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,6 @@

import queue
import threading
import time
from enum import Enum
from typing import Callable, Optional

Expand Down Expand Up @@ -174,17 +173,14 @@ def __init__(
# The state of the command interface must not be changed directly.
# Use the change_state method instead.
self._state = CommandInterfaceState.TERMINATED
self._state_lock = threading.Lock()
# Serializes connect() / disconnect() so the lifecycle transitions
# and the resources they own cannot interleave.
self._lifecycle_lock = threading.Lock()
# Guards _channel, _call and _write_queue, which are replaced by
# the reconnect loop (reader thread) and read/torn down by
# disconnect() and write_request() (caller thread).
self._stream_lock = threading.Lock()
# Guards _state, _channel, _call, _write_queue and _reader_thread.
# Only connect() holds it across a blocking call.
self._lock = threading.Lock()
# Set by disconnect() to stop the reader thread and its reconnects.
self._stop_event = threading.Event()
self._channel: Optional[grpc.Channel] = None
self._call = None
self._write_queue: "queue.Queue" = queue.Queue()
self._write_queue: Optional["queue.Queue"] = None
self._reader_thread: Optional[threading.Thread] = None

@property
Expand All @@ -205,55 +201,72 @@ def connect(self) -> None:
established, a later lost connection is retried
transparently every ``RECONNECT_INTERVAL`` seconds.

:raises ConnectionException: If already connected, or if
the connection could not be established.
:raises ConnectionException: If already connected, if the
previous connection is still shutting down, or if the
connection could not be established.
"""
with self._lifecycle_lock:
if self._state in (
CommandInterfaceState.INITIALIZED,
CommandInterfaceState.CONNECTED,
CommandInterfaceState.RECONNECTING,
):
# Held across the whole attempt (up to CHANNEL_READY_TIMEOUT
# seconds) so a concurrent disconnect() cannot interleave with it.
# No reader thread exists yet that could need the lock.
with self._lock:
if self._state != CommandInterfaceState.TERMINATED:
raise ConnectionException("Already connected.")
# Only one reader thread may ever run: a previous one that did
# not stop in time would otherwise reconnect and replace this
# connection's stream.
if (
self._reader_thread is not None
and self._reader_thread.is_alive()
):
raise ConnectionException(
"Previous connection is still shutting down."
)

# Only change the state once past the point where connecting
# can still fail, so a failed attempt leaves the connection
# exactly as it was and free to retry via a plain connect().
call = self._open_stream()
self.change_state(CommandInterfaceState.INITIALIZED)
channel, call, write_queue = self._open_stream()
self._stop_event.clear()
self._channel = channel
self._call = call
self._write_queue = write_queue
self._set_state(CommandInterfaceState.INITIALIZED)

self._reader_thread = threading.Thread(
target=self._read_from_grpc, args=(call,), daemon=True
)
self._reader_thread.start()
self.change_state(CommandInterfaceState.CONNECTED)
self._set_state(CommandInterfaceState.CONNECTED)
self._logger.debug("Connected to the Ankaios server over gRPC.")

def disconnect(self) -> None:
"""
Disconnect from the gRPC connection.
"""
with self._lifecycle_lock:
with self._lock:
if self._state == CommandInterfaceState.TERMINATED:
self._logger.debug("Already disconnected.")
return

self._logger.debug("Disconnecting..")
self.change_state(CommandInterfaceState.TERMINATED)
with self._stream_lock:
call = self._call
if call is not None:
call.cancel()
if self._reader_thread is not None:
self._reader_thread.join(timeout=2)
if self._reader_thread.is_alive():
self._logger.error("Reader thread did not stop.")
self._reader_thread = None
with self._stream_lock:
if self._channel is not None:
self._channel.close()
self._channel = None
self._call = None
self._set_state(CommandInterfaceState.TERMINATED)
self._stop_event.set()
channel, self._channel = self._channel, None
call, self._call = self._call, None
self._write_queue = None
# Kept even if it does not stop in time, so connect() can
# refuse to start a second reader while it is still alive.
reader_thread = self._reader_thread

# The stream is detached, so the slow teardown runs unlocked.
if call is not None:
call.cancel()
if reader_thread is not None:
reader_thread.join(timeout=2)
if reader_thread.is_alive():
self._logger.error("Reader thread did not stop.")
if channel is not None:
channel.close()

def write_request(self, request: Request) -> None:
"""
Expand All @@ -264,20 +277,16 @@ def write_request(self, request: Request) -> None:

:raises ConnectionException: If not connected.
"""
with self._state_lock:
with self._lock:
if self._state != CommandInterfaceState.CONNECTED:
self._logger.error(
"Could not write to the gRPC connection, not connected."
)
raise ConnectionException(
"Could not write to the gRPC connection, not connected."
)
# Held so the queue cannot be replaced by a concurrent reconnect
# between reading the reference and putting onto it.
with self._stream_lock:
self._write_queue.put(
_grpc_api.ToServer(request=request._to_proto())
)
write_queue = self._write_queue
write_queue.put(_grpc_api.ToServer(request=request._to_proto()))

def _read_from_grpc(self, call) -> None:
"""
Expand All @@ -294,7 +303,7 @@ def _read_from_grpc(self, call) -> None:
for from_server in call:
self._handle_from_server(from_server)
except grpc.RpcError as e:
if self._state == CommandInterfaceState.TERMINATED:
if self._stop_event.is_set():
# disconnect() already cancelled the call itself;
# this is the expected, self-inflicted result.
self._logger.debug(
Expand All @@ -305,13 +314,12 @@ def _read_from_grpc(self, call) -> None:
"Error while reading from the gRPC connection: '%s'",
e,
)
else:
# The stream ended without an error; still need to know
# whether that was disconnect()'s doing before retrying.
if self._state == CommandInterfaceState.TERMINATED:
return

self.change_state(CommandInterfaceState.RECONNECTING)
with self._lock:
# The stream may also have ended because of disconnect().
if self._stop_event.is_set():
return
self._set_state(CommandInterfaceState.RECONNECTING)
self._logger.warning(
"Lost connection to the Ankaios server, attempting to "
"reconnect every %ss..",
Expand Down Expand Up @@ -349,18 +357,35 @@ def _reconnect(self):
called while reconnecting.
:rtype: Optional[grpc.Call]
"""
while True:
time.sleep(self.RECONNECT_INTERVAL)
if self._state == CommandInterfaceState.TERMINATED:
return None
# Waits on the event so a disconnect wakes us immediately.
while not self._stop_event.wait(self.RECONNECT_INTERVAL):
try:
call = self._open_stream()
channel, call, write_queue = self._open_stream()
except ConnectionException as e:
self._logger.debug("Reconnect attempt failed: '%s'", e)
continue
self.change_state(CommandInterfaceState.CONNECTED)

old_channel = None
with self._lock:
# disconnect() may have run while the stream was opening;
# it must not be brought back to life then.
stopped = self._stop_event.is_set()
if not stopped:
old_channel = self._channel
self._channel = channel
self._call = call
self._write_queue = write_queue
self._set_state(CommandInterfaceState.CONNECTED)
if stopped:
call.cancel()
channel.close()
return None
if old_channel is not None:
# The channel of the lost connection.
old_channel.close()
self._logger.info("Reconnected to the Ankaios server.")
return call
return None

def change_state(self, state: CommandInterfaceState) -> None:
"""
Expand All @@ -374,32 +399,46 @@ def change_state(self, state: CommandInterfaceState) -> None:
:param state: The new state.
:type state: CommandInterfaceState
"""
with self._state_lock:
if state == self._state:
self._logger.debug("State is already %s.", state)
return
self._state = state
self._logger.debug("State changed to %s.", state)
with self._lock:
self._set_state(state)

def _set_state(self, state: CommandInterfaceState) -> None:
"""
Change the state of the connection. The caller must hold
the lock.

:param state: The new state.
:type state: CommandInterfaceState
"""
if state == self._state:
self._logger.debug("State is already %s.", state)
return
self._state = state
self._logger.debug("State changed to %s.", state)

def _open_stream(self):
"""
Builds the gRPC channel, sends the initial CommanderHello and
opens the ConnectCommand bidi stream. Used both for the
initial connect and for every reconnect attempt.
initial connect and for every reconnect attempt. The opened
stream is not stored on the connection; the caller does that.

:returns: The opened bidi call, iterable for FromServer
messages.
:rtype: grpc.Call
:returns: The channel, the opened bidi call (iterable for
FromServer messages) and the queue feeding its requests.
:rtype: tuple

:raises ConnectionException: If the channel does not become
ready in time, or the stream could not be opened.
"""
channel = None
try:
channel = self._build_channel()
grpc.channel_ready_future(channel).result(
timeout=self.CHANNEL_READY_TIMEOUT
)
except (grpc.FutureTimeoutError, grpc.RpcError) as e:
if channel is not None:
channel.close()
raise ConnectionException(
f"Could not connect to the Ankaios server: '{e}'"
) from e
Expand All @@ -413,23 +452,19 @@ def _open_stream(self):
)
)
)
call = stub.ConnectCommand(self._request_iterator(write_queue))
return channel, call, write_queue

# Held across the replacement so disconnect() and write_request()
# never observe a half-replaced write_queue/channel/call.
with self._stream_lock:
self._write_queue = write_queue
call = stub.ConnectCommand(self._request_iterator())
self._channel = channel
self._call = call
return call

def _request_iterator(self):
@staticmethod
def _request_iterator(write_queue: "queue.Queue"):
"""
Generator yielding ToServer messages queued via
:func:`write_request` (and the initial hello), until
:func:`disconnect` cancels the call.

:param write_queue: The queue of the stream being fed.
:type write_queue: queue.Queue
"""
write_queue = self._write_queue
while True:
yield write_queue.get()

Expand Down
Loading
Loading