Fix a few concurrency bugs in Citus support (#2710)

- the `_in_fligh` attribute is accessed from multiple threads and must be protected with mutex when it is changed
- allow adding tasks for `_in_fligh.group` from the `sync_pg_dist_node()` method when timeout is reached. Not doing so might result is indefinite transaction if REST API request from worker node failed.
This commit is contained in:
Alexander Kukushkin
2023-06-12 07:52:46 +02:00
committed by GitHub
parent bd951ccdef
commit 2354f8f004
2 changed files with 43 additions and 21 deletions
+37 -19
View File
@@ -76,8 +76,9 @@ class CitusHandler(Thread):
self._connection = Connection() self._connection = Connection()
self._pg_dist_node: Dict[int, PgDistNode] = {} # Cache of pg_dist_node: {groupid: PgDistNode()} self._pg_dist_node: Dict[int, PgDistNode] = {} # Cache of pg_dist_node: {groupid: PgDistNode()}
self._tasks: List[PgDistNode] = [] # Requests to change pg_dist_node, every task is a `PgDistNode` self._tasks: List[PgDistNode] = [] # Requests to change pg_dist_node, every task is a `PgDistNode`
self._condition = Condition() # protects _pg_dist_node, _tasks, and _schedule_load_pg_dist_node
self._in_flight: Optional[PgDistNode] = None # Reference to the `PgDistNode` being changed in a transaction self._in_flight: Optional[PgDistNode] = None # Reference to the `PgDistNode` being changed in a transaction
self._schedule_load_pg_dist_node = True # Flag that "pg_dist_node" should be queried from the database
self._condition = Condition() # protects _pg_dist_node, _tasks, _in_flight, and _schedule_load_pg_dist_node
self.schedule_cache_rebuild() self.schedule_cache_rebuild()
def is_enabled(self) -> bool: def is_enabled(self) -> bool:
@@ -117,7 +118,8 @@ class CitusHandler(Thread):
except Exception as e: except Exception as e:
logger.error('Exception when executing query "%s", (%s): %r', sql, params, e) logger.error('Exception when executing query "%s", (%s): %r', sql, params, e)
self._connection.close() self._connection.close()
self._in_flight = None with self._condition:
self._in_flight = None
self.schedule_cache_rebuild() self.schedule_cache_rebuild()
raise e raise e
@@ -215,17 +217,20 @@ class CitusHandler(Thread):
def process_task(self, task: PgDistNode) -> bool: def process_task(self, task: PgDistNode) -> bool:
"""Updates a single row in `pg_dist_node` table, optionally in a transaction. """Updates a single row in `pg_dist_node` table, optionally in a transaction.
The transaction is started if we do a demote of the worker node The transaction is started if we do a demote of the worker node or before promoting the other worker if
or before promoting the other worker if there is not transaction there is no transaction in progress. And, the transaction is committed when the switchover/failover completed.
in progress. And, the transaction it is committed when the
switchover/failover completed.
This method returns `True` if node was updated (optionally, .. note:
transaction was committed) as an indicator that The maximum lifetime of the transaction in progress is controlled outside of this method.
the `self._pg_dist_node` cache should be updated.
The maximum lifetime of the transaction in progress .. note:
is controlled outside of this method.""" Read access to `self._in_flight` isn't protected because we know it can't be changed outside of our thread.
:param task: reference to a :class:`PgDistNode` object that represents a row to be updated/created.
:returns: `True` if the row was succesfully created/updated or transaction in progress
was committed as an indicator that the `self._pg_dist_node` cache should be updated,
or, if the new transaction was opened, this method returns `False`.
"""
if task.event == 'after_promote': if task.event == 'after_promote':
# The after_promote may happen without previous before_demote and/or # The after_promote may happen without previous before_demote and/or
@@ -236,7 +241,6 @@ class CitusHandler(Thread):
self.update_node(task) self.update_node(task)
if self._in_flight: if self._in_flight:
self.query('COMMIT') self.query('COMMIT')
self._in_flight = None
return True return True
else: # before_demote, before_promote else: # before_demote, before_promote
if task.timeout: if task.timeout:
@@ -244,11 +248,11 @@ class CitusHandler(Thread):
if not self._in_flight: if not self._in_flight:
self.query('BEGIN') self.query('BEGIN')
self.update_node(task) self.update_node(task)
self._in_flight = task
return False return False
def process_tasks(self) -> None: def process_tasks(self) -> None:
while True: while True:
# Read access to `_in_flight` isn't protected because we know it can't be changed outside of our thread.
if not self._in_flight and not self.load_pg_dist_node(): if not self._in_flight and not self.load_pg_dist_node():
break break
@@ -259,11 +263,17 @@ class CitusHandler(Thread):
update_cache = self.process_task(task) update_cache = self.process_task(task)
except Exception as e: except Exception as e:
logger.error('Exception when working with pg_dist_node: %r', e) logger.error('Exception when working with pg_dist_node: %r', e)
update_cache = False update_cache = None
with self._condition: with self._condition:
if self._tasks: if self._tasks:
if update_cache: if update_cache:
self._pg_dist_node[task.group] = task self._pg_dist_node[task.group] = task
if update_cache is False: # an indicator that process_tasks has started a transaction
self._in_flight = task
else:
self._in_flight = None
if id(self._tasks[i]) == id(task): if id(self._tasks[i]) == id(task):
self._tasks.pop(i) self._tasks.pop(i)
task.wakeup() task.wakeup()
@@ -293,11 +303,19 @@ class CitusHandler(Thread):
with self._condition: with self._condition:
i = self.find_task_by_group(task.group) i = self.find_task_by_group(task.group)
# task.timeout is None is an indicator that it was scheduled # The `PgDistNode.timeout` == None is an indicator that it was scheduled from the sync_pg_dist_node().
# from the sync_pg_dist_node() and we don't want to override if task.timeout is None:
# already existing task created from REST API. # We don't want to override the already existing task created from REST API.
if task.timeout is None and (i is not None or self._in_flight and self._in_flight.group == task.group): if i is not None and self._tasks[i].timeout is not None:
return False return False
# There is a little race condition with tasks created from REST API - the call made "before" the member
# key is updated in DCS. Therefore it is possible that :func:`sync_pg_dist_node` will try to create a
# task based on the outdated values of "state"/"role". To solve it we introduce an artificial timeout.
# Only when the timeout is reached new tasks could be scheduled from sync_pg_dist_node()
if self._in_flight and self._in_flight.group == task.group and self._in_flight.timeout is not None\
and self._in_flight.deadline > time.time():
return False
# Override already existing task for the same worker group # Override already existing task for the same worker group
if i is not None: if i is not None:
+6 -2
View File
@@ -1,3 +1,4 @@
import time
from mock import Mock, patch from mock import Mock, patch
from patroni.postgresql.citus import CitusHandler from patroni.postgresql.citus import CitusHandler
@@ -16,7 +17,7 @@ class TestCitus(BaseTestPostgresql):
self.cluster = get_cluster_initialized_with_leader() self.cluster = get_cluster_initialized_with_leader()
self.cluster.workers[1] = self.cluster self.cluster.workers[1] = self.cluster
@patch('time.time', Mock(side_effect=[100, 130, 160, 190, 220, 250, 280, 310])) @patch('time.time', Mock(side_effect=[100, 130, 160, 190, 220, 250, 280, 310, 340, 370]))
@patch('patroni.postgresql.citus.logger.exception', Mock(side_effect=SleepException)) @patch('patroni.postgresql.citus.logger.exception', Mock(side_effect=SleepException))
@patch('patroni.postgresql.citus.logger.warning') @patch('patroni.postgresql.citus.logger.warning')
@patch('patroni.postgresql.citus.PgDistNode.wait', Mock()) @patch('patroni.postgresql.citus.PgDistNode.wait', Mock())
@@ -66,11 +67,14 @@ class TestCitus(BaseTestPostgresql):
mock_logger.assert_called_once() mock_logger.assert_called_once()
self.assertTrue(mock_logger.call_args[0][0].startswith('Overriding existing task:')) self.assertTrue(mock_logger.call_args[0][0].startswith('Overriding existing task:'))
# add_task called from sync_pg_dist_node should not override already scheduled or in flight task # add_task called from sync_pg_dist_node should not override already scheduled or in flight task until deadline
self.assertIsNotNone(self.c.add_task('after_promote', 1, 'postgres://host:5432/postgres', 30)) self.assertIsNotNone(self.c.add_task('after_promote', 1, 'postgres://host:5432/postgres', 30))
self.assertIsNone(self.c.add_task('after_promote', 1, 'postgres://host:5432/postgres')) self.assertIsNone(self.c.add_task('after_promote', 1, 'postgres://host:5432/postgres'))
self.c._in_flight = self.c._tasks.pop() self.c._in_flight = self.c._tasks.pop()
self.c._in_flight.deadline = self.c._in_flight.timeout + time.time()
self.assertIsNone(self.c.add_task('after_promote', 1, 'postgres://host:5432/postgres')) self.assertIsNone(self.c.add_task('after_promote', 1, 'postgres://host:5432/postgres'))
self.c._in_flight.deadline = 0
self.assertIsNotNone(self.c.add_task('after_promote', 1, 'postgres://host:5432/postgres'))
# If there is no transaction in progress and cached pg_dist_node matching desired state task should not be added # If there is no transaction in progress and cached pg_dist_node matching desired state task should not be added
self.c._schedule_load_pg_dist_node = False self.c._schedule_load_pg_dist_node = False