diff --git a/patroni/ctl.py b/patroni/ctl.py index 7fc4c23a..6b899f3d 100644 --- a/patroni/ctl.py +++ b/patroni/ctl.py @@ -268,7 +268,7 @@ def get_cursor(cluster, connect_parameters, role='master', member=None): return None -def get_members(cluster, cluster_name, member_names, role, force, action): +def get_members(cluster, cluster_name, member_names, role, force, action, scheduled_at=None): candidates = {m.name: m for m in cluster.members} if not force or role: @@ -291,10 +291,18 @@ def get_members(cluster, cluster_name, member_names, role, force, action): if member_name not in candidates: raise PatroniCtlException('{0} is not a member of cluster'.format(member_name)) - if not force: - confirm = click.confirm('Are you sure you want to {0} members {1}?'.format(action, ', '.join(member_names))) - if not confirm: - raise PatroniCtlException('Aborted {0}'.format(action)) + if scheduled_at: + if not force: + confirm = click.confirm('Are you sure you want to schedule {0} of members {1} at {2}?' + .format(action, ', '.join(member_names), scheduled_at)) + if not confirm: + raise PatroniCtlException('Aborted scheduled {0}'.format(action)) + else: + if not force: + confirm = click.confirm('Are you sure you want to {0} members {1}?' + .format(action, ', '.join(member_names))) + if not confirm: + raise PatroniCtlException('Aborted {0}'.format(action)) return [candidates[n] for n in member_names] @@ -515,7 +523,14 @@ def reload(obj, cluster_name, member_names, force, role): def restart(obj, cluster_name, member_names, force, role, p_any, scheduled, version, pending, timeout): cluster = get_dcs(obj, cluster_name).get_cluster() - members = get_members(cluster, cluster_name, member_names, role, force, 'restart') + if scheduled is None and not force: + next_hour = (datetime.datetime.now() + datetime.timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M') + scheduled = click.prompt('When should the restart take place (e.g. ' + next_hour + ') ', + type=str, default='now') + + scheduled_at = parse_scheduled(scheduled) + + members = get_members(cluster, cluster_name, member_names, role, force, 'restart', scheduled_at) if p_any: random.shuffle(members) members = members[:1] @@ -536,10 +551,6 @@ def restart(obj, cluster_name, member_names, force, role, p_any, scheduled, vers content['postgres_version'] = version - if scheduled is None and not force: - scheduled = click.prompt('When should the restart take place (e.g. 2015-10-01T14:30) ', type=str, default='now') - - scheduled_at = parse_scheduled(scheduled) if scheduled_at: if cluster.is_paused(): raise PatroniCtlException("Can't schedule restart in the paused state") @@ -635,7 +646,8 @@ def _do_failover_or_switchover(obj, action, cluster_name, master, candidate, for if action == 'switchover': if scheduled is None and not force: - scheduled = click.prompt('When should the switchover take place (e.g. 2015-10-01T14:30) ', + next_hour = (datetime.datetime.now() + datetime.timedelta(hours=1)).strftime('%Y-%m-%dT%H:%M') + scheduled = click.prompt('When should the switchover take place (e.g. ' + next_hour + ' ) ', type=str, default='now') scheduled_at = parse_scheduled(scheduled) @@ -654,9 +666,14 @@ def _do_failover_or_switchover(obj, action, cluster_name, master, candidate, for if not force: demote_msg = ', demoting current master ' + master if master else '' - - if not click.confirm('Are you sure you want to {0} cluster {1}{2}?'.format(action, cluster_name, demote_msg)): - raise PatroniCtlException('Aborting ' + action) + if scheduled_at_str: + if not click.confirm('Are you sure you want to schedule {0} of cluster {1} at {2}{3}?' + .format(action, cluster_name, scheduled_at_str, demote_msg)): + raise PatroniCtlException('Aborting scheduled ' + action) + else: + if not click.confirm('Are you sure you want to {0} cluster {1}{2}?' + .format(action, cluster_name, demote_msg)): + raise PatroniCtlException('Aborting ' + action) r = None try: diff --git a/tests/test_ctl.py b/tests/test_ctl.py index e4e47e94..db04cbf6 100644 --- a/tests/test_ctl.py +++ b/tests/test_ctl.py @@ -90,10 +90,14 @@ class TestCtl(unittest.TestCase): result = self.runner.invoke(ctl, ['switchover', 'dummy', '--force', '--scheduled', '2015-01-01T12:00:00']) assert result.exit_code == 1 - # Aborting switchover, as we anser NO to the confirmation + # Aborting switchover, as we answer NO to the confirmation result = self.runner.invoke(ctl, ['switchover', 'dummy'], input='leader\nother\n\nN') assert result.exit_code == 1 + # Aborting scheduled switchover, as we answer NO to the confirmation + result = self.runner.invoke(ctl, ['switchover', 'dummy', '--scheduled', '2015-01-01T12:00:00+01:00'], input='leader\nother\n\nN') + assert result.exit_code == 1 + # Target and source are equal result = self.runner.invoke(ctl, ['switchover', 'dummy'], input='leader\nleader\n\ny') assert result.exit_code == 1 @@ -246,7 +250,7 @@ class TestCtl(unittest.TestCase): @patch('patroni.ctl.get_dcs') def test_restart_reinit(self, mock_get_dcs): mock_get_dcs.return_value.get_cluster = get_cluster_initialized_with_leader - result = self.runner.invoke(ctl, ['restart', 'alpha'], input='y\n\nnow') + result = self.runner.invoke(ctl, ['restart', 'alpha'], input='now\ny\n') assert 'Failed: restart for' in result.output assert result.exit_code == 0 @@ -258,18 +262,22 @@ class TestCtl(unittest.TestCase): assert result.exit_code == 0 # Aborted restart - result = self.runner.invoke(ctl, ['restart', 'alpha'], input='N') + result = self.runner.invoke(ctl, ['restart', 'alpha'], input='now\nN') assert result.exit_code == 1 result = self.runner.invoke(ctl, ['restart', 'alpha', '--pending', '--force']) assert result.exit_code == 0 + # Aborted scheduled restart + result = self.runner.invoke(ctl, ['restart', 'alpha', '--scheduled', '2019-10-01T14:30'], input='N') + assert result.exit_code == 1 + # Not a member - result = self.runner.invoke(ctl, ['restart', 'alpha', 'dummy', '--any'], input='y') + result = self.runner.invoke(ctl, ['restart', 'alpha', 'dummy', '--any'], input='now\ny') assert result.exit_code == 1 # Wrong pg version - result = self.runner.invoke(ctl, ['restart', 'alpha', '--any', '--pg-version', '9.1'], input='y') + result = self.runner.invoke(ctl, ['restart', 'alpha', '--any', '--pg-version', '9.1'], input='now\ny') assert 'Error: Invalid PostgreSQL version format' in result.output assert result.exit_code == 1