diff --git a/docs/replica_bootstrap.rst b/docs/replica_bootstrap.rst index 861a0e86..a6e3db8e 100644 --- a/docs/replica_bootstrap.rst +++ b/docs/replica_bootstrap.rst @@ -71,6 +71,22 @@ Makes the configured ``command`` to be called additionally with ``--arg1=value1 .. note:: Bootstrap methods are neither chained, nor fallen-back to the default one in case the primary one fails +As an example, you are able to bootstrap a fresh Patroni cluster from a Barman backup with a configuration like this: + +.. code:: YAML + + bootstrap: + method: barman + barman: + keep_existing_recovery_conf: true + command: patroni_barman_recover + api-url: https://barman-host:7480 + barman-server: my_server + ssh-command: ssh postgres@patroni-host + +.. note:: + ``patroni_barman_recover`` requires that you have both Barman and ``pg-backup-api`` configured in the Barman host, so it can execute a remote ``barman recover`` through the backup API. + The above example uses a subset of the available parameters. You can get more information running ``patroni_barman_recover --help``. .. _custom_replica_creation: @@ -125,6 +141,25 @@ example: pgbackrest basebackup: max-rate: '100M' +example: Barman + +.. code:: YAML + + postgresql: + create_replica_methods: + - barman + - basebackup + barman: + command: patroni_barman_recover + api-url: https://barman-host:7480 + barman-server: my_server + ssh-command: ssh postgres@patroni-host + basebackup: + max-rate: '100M' + +.. note:: + ``patroni_barman_recover`` requires that you have both Barman and ``pg-backup-api`` configured in the Barman host, so it can execute a remote ``barman recover`` through the backup API. + The above example uses a subset of the available parameters. You can get more information running ``patroni_barman_recover --help``. The ``create_replica_methods`` defines available replica creation methods and the order of executing them. Patroni will stop on the first one that returns 0. Each method should define a separate section in the configuration file, listing the command diff --git a/patroni/scripts/barman_recover.py b/patroni/scripts/barman_recover.py new file mode 100644 index 00000000..1cffe34d --- /dev/null +++ b/patroni/scripts/barman_recover.py @@ -0,0 +1,468 @@ +#!/usr/bin/env python + +"""Restore a Barman backup to the local node through ``pg-backup-api``. + +This script can be used both as a custom bootstrap method, and as a custom +create replica method. Check the output of ``--help`` to understand the +parameters supported by the script. ``--datadir`` is a special parameter and it +is automatically filled by Patroni in both cases. + +It requires that you have previously configured a Barman server, and that you +have ``pg-backup-api`` configured and running in the same host as Barman. + +Refer to :class:`ExitCode` for possible exit codes of this script. +""" +from argparse import ArgumentParser +from enum import IntEnum +import json +import logging +import sys +import time +from typing import Any, Callable, Optional, Tuple, Type, Union +from urllib.parse import urljoin +from urllib3 import PoolManager +from urllib3.exceptions import MaxRetryError +from urllib3.response import HTTPResponse + + +class ExitCode(IntEnum): + """Possible exit codes of this script. + + :cvar RECOVERY_DONE: backup was successfully restored. + :cvar RECOVERY_FAILED: recovery of the backup faced an issue. + :cvar API_NOT_OK: ``pg-backup-api`` status is not ``OK``. + :cvar HTTP_REQUEST_ERROR: an error has occurred during a request to the + ``pg-backup-api``. + :cvar HTTP_RESPONSE_MALFORMED: ``pg-backup-api`` returned a bogus response. + """ + + RECOVERY_DONE = 0 + RECOVERY_FAILED = 1 + API_NOT_OK = 2 + HTTP_REQUEST_ERROR = 3 + HTTP_RESPONSE_MALFORMED = 4 + + +class RetriesExceeded(Exception): + """Maximum number of retries exceeded.""" + + +def retry(exceptions: Union[Type[Exception], Tuple[Type[Exception], ...]]) \ + -> Any: + """Retry an operation n times if expected *exceptions* are faced. + + .. note:: + Should be used as a decorator of a class' method as it expects the + first argument to be a class instance. + + The class which method is going to be decorated should contain a couple + attributes: + + * ``max_retries``: maximum retry attempts before failing; + * ``retry_wait``: how long to wait before retrying. + + :param exceptions: exceptions that could trigger a retry attempt. + + :raises: + :exc:`RetriesExceeded`: if the maximum number of attempts has been + exhausted. + """ + def decorator(func: Callable[..., Any]) -> Any: + def inner_func(instance: object, *args: Any, **kwargs: Any) -> Any: + times: int = getattr(instance, "max_retries") + retry_wait: int = getattr(instance, "retry_wait") + method_name = f"{instance.__class__.__name__}.{func.__name__}" + + attempt = 1 + + while attempt <= times: + try: + return func(instance, *args, **kwargs) + except exceptions as exc: + logging.warning("Attempt %d of %d on method %s failed " + "with %r.", + attempt, times, method_name, exc) + attempt += 1 + + time.sleep(retry_wait) + + raise RetriesExceeded("Maximum number of retries exceeded for " + f"method {method_name}.") + return inner_func + return decorator + + +class BarmanRecover: + """Facilities for performing a remote ``barman recover`` operation. + + You should instantiate this class, which will take care of configuring the + operation accordingly. When you want to start the operation, you should + call :meth:`restore_backup`. At any point of interaction with this class, + you may face a :func:`sys.exit` call. Refer to :class:`ExitCode` for a view + on the possible exit codes. + + :ivar api_url: base URL to reach the ``pg-backup-api``. + :ivar cert_file: certificate to authenticate against the + ``pg-backup-api``, if required. + :ivar key_file: certificate key to authenticate against the + ``pg-backup-api``, if required. + :ivar barman_server: name of the Barman server which backup is to be + restored. + :ivar backup_id: ID of the backup from the Barman server. + :ivar ssh_command: SSH command to connect from the Barman host to the + local host. + :ivar data_directory: path to the Postgres data directory where to + restore the backup at. + :ivar loop_wait: how long to wait before checking again the status of the + recovery process. Higher values are useful for backups that are + expected to take long to restore. + :ivar retry_wait: how long to wait before retrying a failed request to the + ``pg-backup-api``. + :ivar max_retries: maximum number of retries when ``pg-backup-api`` returns + malformed responses. + :ivar http: a HTTP pool manager for performing web requests. + """ + + def __init__(self, api_url: str, barman_server: str, backup_id: str, + ssh_command: str, data_directory: str, loop_wait: int, + retry_wait: int, max_retries: int, + cert_file: Optional[str] = None, + key_file: Optional[str] = None) -> None: + """Create a new instance of :class:`BarmanRecover`. + + Make sure the ``pg-backup-api`` is reachable and running fine. + + :param api_url: base URL to reach the ``pg-backup-api``. + :param barman_server: name of the Barman server which backup is to be + restored. + :param backup_id: ID of the backup from the Barman server. + :param ssh_command: SSH command to connect from the Barman host to the + local host. + :param data_directory: path to the Postgres data directory where to + restore the backup at. + :param loop_wait: how long to wait before checking again the status of + the recovery process. Higher values are useful for backups that are + expected to take long to restore. + :param retry_wait: how long to wait before retrying a failed request to + the ``pg-backup-api``. + :param max_retries: maximum number of retries when ``pg-backup-api`` + returns malformed responses. + :param cert_file: certificate to authenticate against the + ``pg-backup-api``, if required. + :param key_file: certificate key to authenticate against the + ``pg-backup-api``, if required. + """ + self.api_url = api_url + self.cert_file = cert_file + self.key_file = key_file + self.barman_server = barman_server + self.backup_id = backup_id + self.ssh_command = ssh_command + self.data_directory = data_directory + self.loop_wait = loop_wait + self.retry_wait = retry_wait + self.max_retries = max_retries + self.http = PoolManager(cert_file=cert_file, key_file=key_file) + self._ensure_api_ok() + + def _build_full_url(self, url_path: str) -> str: + """Build the full URL by concatenating *url_path* with the base URL. + + :param url_path: path to be accessed in the ``pg-backup-api``. + + :returns: the full URL after concatenating. + """ + return urljoin(self.api_url, url_path) + + @staticmethod + def _deserialize_response(response: HTTPResponse) -> Any: + """Retrieve body from *response* as a deserialized JSON object. + + :param response: response from which JSON body will be deserialized. + + :returns: the deserialized JSON body. + """ + return json.loads(response.data.decode("utf-8")) + + @staticmethod + def _serialize_request(body: Any) -> Any: + """Serialize a request body. + + :param body: content of the request body to be serialized. + + :returns: the serialized request body. + """ + return json.dumps(body).encode("utf-8") + + def _get_request(self, url_path: str) -> Any: + """Perform a ``GET`` request to *url_path*. + + .. note:: + If a :exc:`MaxRetryError` is faced while performing the request, + then exit with :attr:`ExitCode.HTTP_REQUEST_ERROR` + + :param url_path: URL to perform the ``GET`` request against. + + :returns: the deserialized response body. + """ + response = None + + try: + response = self.http.request("GET", self._build_full_url(url_path)) + except MaxRetryError as exc: + logging.critical("An error occurred while performing an HTTP GET " + "request: %r", exc) + sys.exit(ExitCode.HTTP_REQUEST_ERROR) + + return self._deserialize_response(response) + + def _post_request(self, url_path: str, body: Any) -> Any: + """Perform a ``POST`` request to *url_path* serializing *body* as JSON. + + .. note:: + If a :exc:`MaxRetryError` is faced while performing the request, + then exit with :attr:`ExitCode.HTTP_REQUEST_ERROR` + + :param url_path: URL to perform the ``POST`` request against. + :param body: the body to be serialized as JSON and sent in the request. + + :returns: the deserialized response body. + """ + body = self._serialize_request(body) + + response = None + + try: + response = self.http.request("POST", + self._build_full_url(url_path), + body=body, + headers={ + "Content-Type": "application/json" + }) + except MaxRetryError as exc: + logging.critical("An error occurred while performing an HTTP POST " + "request: %r", exc) + sys.exit(ExitCode.HTTP_REQUEST_ERROR) + + return self._deserialize_response(response) + + def _ensure_api_ok(self) -> None: + """Ensure ``pg-backup-api`` is reachable and ``OK``. + + .. note:: + If ``pg-backup-api`` status is not ``OK``, then exit with + :attr:`ExitCode.API_NOT_OK`. + """ + response = self._get_request("status") + + if response != "OK": + logging.critical("pg-backup-api is not working: %s", response) + sys.exit(ExitCode.API_NOT_OK) + + @retry(KeyError) + def _create_recovery_operation(self) -> str: + """Create a recovery operation on the ``pg-backup-api``. + + :returns: the ID of the recovery operation that has been created. + """ + response = self._post_request( + f"servers/{self.barman_server}/operations", + { + "type": "recovery", + "backup_id": self.backup_id, + "remote_ssh_command": self.ssh_command, + "destination_directory": self.data_directory, + }, + ) + + return response["operation_id"] + + @retry(KeyError) + def _get_recovery_operation_status(self, operation_id: str) -> str: + """Get status of the recovery operation *operation_id*. + + :param operation_id: ID of the recovery operation to be checked. + + :returns: the status of the recovery operation. + """ + response = self._get_request( + f"servers/{self.barman_server}/operations/{operation_id}", + ) + + return response["status"] + + def restore_backup(self) -> bool: + """Restore the configured Barman backup through ``pg-backup-api``. + + .. note:: + If recovery API request returns a malformed response, then exit with + :attr:`ExitCode.HTTP_RESPONSE_MALFORMED`. + + :returns: ``True`` if it was successfully recovered, ``False`` + otherwise. + """ + operation_id = None + + try: + operation_id = self._create_recovery_operation() + except RetriesExceeded: + logging.critical("Maximum number of retries exceeded, exiting.") + sys.exit(ExitCode.HTTP_RESPONSE_MALFORMED) + + logging.info("Created the recovery operation with ID %s", operation_id) + + status = None + + while True: + try: + status = self._get_recovery_operation_status(operation_id) + except RetriesExceeded: + logging.critical("Maximum number of retries exceeded, " + "exiting.") + sys.exit(ExitCode.HTTP_RESPONSE_MALFORMED) + + if status != "IN_PROGRESS": + break + + logging.info("Recovery operation %s is still in progress", + operation_id) + time.sleep(self.loop_wait) + + return status == "DONE" + + +def set_up_logging(log_file: Optional[str] = None) -> None: + """Set up logging to file, if *log_file* is given, otherwise to console. + + :param log_file: file where to log messages, if any. + """ + logging.basicConfig(filename=log_file, level=logging.INFO, + format="%(asctime)s %(levelname)s: %(message)s") + + +def main() -> None: + """Entry point of this script. + + Parse the command-line arguments and recover a Barman backup through + ``pg-backup-api`` to the local host. + """ + parser = ArgumentParser( + epilog=( + "Wrapper script for ``pg-backup-api``. Communicate with the API " + "running at ``--api-url`` to restore a ``--backup-id`` Barman " + "backup of the server ``--barman-server``." + ), + ) + parser.add_argument( + "--api-url", + type=str, + required=True, + help="URL to reach the ``pg-backup-api``, e.g. " + "``http://localhost:7480``", + dest="api_url", + ) + parser.add_argument( + "--cert-file", + type=str, + required=False, + help="Certificate to authenticate against the API, if required.", + dest="cert_file", + ) + parser.add_argument( + "--key-file", + type=str, + required=False, + help="Certificate key to authenticate against the API, if required.", + dest="key_file", + ) + parser.add_argument( + "--barman-server", + type=str, + required=True, + help="Name of the Barman server from which to restore the backup.", + dest="barman_server", + ) + parser.add_argument( + "--backup-id", + type=str, + required=False, + default="latest", + help="ID of the Barman backup to be restored. You can use any value " + "supported by ``barman recover`` command " + "(default: ``%(default)s``)", + dest="backup_id", + ) + parser.add_argument( + "--ssh-command", + type=str, + required=True, + help="Value to be passed as ``--remote-ssh-command`` to " + "``barman recover``.", + dest="ssh_command", + ) + parser.add_argument( + "--data-directory", + "--datadir", + type=str, + required=True, + help="Destination path where to restore the barman backup in the " + "local host.", + dest="data_directory", + ) + parser.add_argument( + "--log-file", + type=str, + required=False, + help="File where to log messages produced by this script, if any.", + dest="log_file", + ) + parser.add_argument( + "--loop-wait", + type=int, + required=False, + default=10, + help="How long to wait before checking again the status of the " + "recovery process, in seconds. Use higher values if your " + "recovery is expected to take long (default: ``%(default)s``)", + dest="loop_wait", + ) + parser.add_argument( + "--retry-wait", + type=int, + required=False, + default=2, + help="How long to wait before retrying a failed ``pg-backup-api`` " + "request (default: ``%(default)s``)", + dest="retry_wait", + ) + parser.add_argument( + "--max-retries", + type=int, + required=False, + default=5, + help="Maximum number of retries when receiving malformed responses " + "from the ``pg-backup-api`` (default: ``%(default)s``)", + dest="max_retries", + ) + args, _ = parser.parse_known_args() + + set_up_logging(args.log_file) + + barman_recover = BarmanRecover(args.api_url, args.barman_server, + args.backup_id, args.ssh_command, + args.data_directory, args.loop_wait, + args.retry_wait, args.max_retries, + args.cert_file, args.key_file) + + successful = barman_recover.restore_backup() + + if successful: + logging.info("Recovery operation finished successfully.") + sys.exit(ExitCode.RECOVERY_DONE) + else: + logging.critical("Recovery operation failed.") + sys.exit(ExitCode.RECOVERY_FAILED) + + +if __name__ == "__main__": + main() diff --git a/setup.py b/setup.py index 4c1c25b3..5bd9c71f 100644 --- a/setup.py +++ b/setup.py @@ -54,7 +54,8 @@ CONSOLE_SCRIPTS = ['patroni = patroni.__main__:main', 'patronictl = patroni.ctl:ctl', 'patroni_raft_controller = patroni.raft_controller:main', "patroni_wale_restore = patroni.scripts.wale_restore:main", - "patroni_aws = patroni.scripts.aws:main"] + "patroni_aws = patroni.scripts.aws:main", + "patroni_barman_recover = patroni.scripts.barman_recover:main"] class _Command(Command): diff --git a/tests/test_barman_recover.py b/tests/test_barman_recover.py new file mode 100644 index 00000000..4334d6ae --- /dev/null +++ b/tests/test_barman_recover.py @@ -0,0 +1,364 @@ +import logging +from mock import MagicMock, Mock, call, patch +import unittest +from urllib3.exceptions import MaxRetryError + +from patroni.scripts.barman_recover import BarmanRecover, ExitCode, RetriesExceeded, main, set_up_logging + + +API_URL = "http://localhost:7480" +BARMAN_SERVER = "my_server" +BACKUP_ID = "backup_id" +SSH_COMMAND = "ssh postgres@localhost" +DATA_DIRECTORY = "/path/to/pgdata" +LOOP_WAIT = 10 +RETRY_WAIT = 2 +MAX_RETRIES = 5 + + +class TestBarmanRecover(unittest.TestCase): + + @patch.object(BarmanRecover, "_ensure_api_ok", Mock()) + @patch("patroni.scripts.barman_recover.PoolManager", MagicMock()) + def setUp(self): + self.br = BarmanRecover(API_URL, BARMAN_SERVER, BACKUP_ID, SSH_COMMAND, DATA_DIRECTORY, LOOP_WAIT, RETRY_WAIT, + MAX_RETRIES) + # Reset the mock as the same instance is used across tests + self.br.http.request.reset_mock() + self.br.http.request.side_effect = None + + def test__build_full_url(self): + self.assertEqual(self.br._build_full_url("/some/path"), f"{API_URL}/some/path") + + @patch("json.loads") + def test__deserialize_response(self, mock_json_loads): + mock_response = MagicMock() + self.assertIsNotNone(self.br._deserialize_response(mock_response)) + mock_json_loads.assert_called_once_with(mock_response.data.decode("utf-8")) + + @patch("json.dumps") + def test__serialize_request(self, mock_json_dumps): + body = "some_body" + ret = self.br._serialize_request(body) + self.assertIsNotNone(ret) + mock_json_dumps.assert_called_once_with(body) + mock_json_dumps.return_value.encode.assert_called_once_with("utf-8") + + @patch.object(BarmanRecover, "_deserialize_response", Mock(return_value="test")) + @patch("logging.critical") + def test__get_request(self, mock_logging): + mock_request = self.br.http.request + + # with no error + self.assertEqual(self.br._get_request("/some/path"), "test") + mock_request.assert_called_once_with("GET", f"{API_URL}/some/path") + + # with MaxRetryError + http_error = MaxRetryError(self.br.http, f"{API_URL}/some/path") + mock_request.side_effect = http_error + + with self.assertRaises(SystemExit) as exc: + self.assertIsNone(self.br._get_request("/some/path")) + + mock_logging.assert_called_once_with("An error occurred while performing an HTTP GET request: %r", http_error) + self.assertEqual(exc.exception.code, ExitCode.HTTP_REQUEST_ERROR) + + # with Exception + mock_logging.reset_mock() + mock_request.side_effect = Exception("Some error.") + + with patch("sys.exit") as mock_sys: + with self.assertRaises(Exception): + self.assertIsNone(self.br._get_request("/some/path")) + + mock_logging.assert_not_called() + mock_sys.assert_not_called() + + @patch.object(BarmanRecover, "_deserialize_response", Mock(return_value="test")) + @patch("logging.critical") + @patch.object(BarmanRecover, "_serialize_request") + def test__post_request(self, mock_serialize, mock_logging): + mock_request = self.br.http.request + + # with no error + self.assertEqual(self.br._post_request("/some/path", "some body"), "test") + mock_serialize.assert_called_once_with("some body") + mock_request.assert_called_once_with("POST", f"{API_URL}/some/path", body=mock_serialize.return_value, + headers={"Content-Type": "application/json"}) + + # with HTTPError + http_error = MaxRetryError(self.br.http, f"{API_URL}/some/path") + mock_request.side_effect = http_error + + with self.assertRaises(SystemExit) as exc: + self.assertIsNone(self.br._post_request("/some/path", "some body")) + + mock_logging.assert_called_once_with("An error occurred while performing an HTTP POST request: %r", http_error) + self.assertEqual(exc.exception.code, ExitCode.HTTP_REQUEST_ERROR) + + # with Exception + mock_logging.reset_mock() + mock_request.side_effect = Exception("Some error.") + + with patch("sys.exit") as mock_sys: + with self.assertRaises(Exception): + self.br._post_request("/some/path", "some body") + + mock_logging.assert_not_called() + mock_sys.assert_not_called() + + @patch("logging.critical") + @patch.object(BarmanRecover, "_get_request") + def test__ensure_api_ok(self, mock_get_request, mock_logging): + # API ok + mock_get_request.return_value = "OK" + + with patch("sys.exit") as mock_sys: + self.assertIsNone(self.br._ensure_api_ok()) + mock_logging.assert_not_called() + mock_sys.assert_not_called() + + # API not ok + mock_get_request.return_value = "random" + + with self.assertRaises(SystemExit) as exc: + self.assertIsNone(self.br._ensure_api_ok()) + + mock_logging.assert_called_once_with("pg-backup-api is not working: %s", "random") + self.assertEqual(exc.exception.code, ExitCode.API_NOT_OK) + + @patch("logging.warning") + @patch("time.sleep") + @patch.object(BarmanRecover, "_post_request") + def test__create_recovery_operation(self, mock_post_request, mock_sleep, mock_logging): + # well formed response + mock_post_request.return_value = {"operation_id": "some_id"} + self.assertEqual(self.br._create_recovery_operation(), "some_id") + mock_sleep.assert_not_called() + mock_logging.assert_not_called() + mock_post_request.assert_called_once_with( + f"servers/{BARMAN_SERVER}/operations", + { + "type": "recovery", + "backup_id": BACKUP_ID, + "remote_ssh_command": SSH_COMMAND, + "destination_directory": DATA_DIRECTORY, + } + ) + + # malformed response + mock_post_request.return_value = {"operation_idd": "some_id"} + + with self.assertRaises(RetriesExceeded) as exc: + self.br._create_recovery_operation() + + self.assertEqual(str(exc.exception), + "Maximum number of retries exceeded for method BarmanRecover._create_recovery_operation.") + + self.assertEqual(mock_sleep.call_count, self.br.max_retries) + mock_sleep.assert_has_calls([call(self.br.retry_wait)] * self.br.max_retries) + + self.assertEqual(mock_logging.call_count, self.br.max_retries) + for i in range(mock_logging.call_count): + call_args = mock_logging.mock_calls[i].args + self.assertEqual(len(call_args), 5) + self.assertEqual(call_args[0], "Attempt %d of %d on method %s failed with %r.") + self.assertEqual(call_args[1], i + 1) + self.assertEqual(call_args[2], self.br.max_retries) + self.assertEqual(call_args[3], "BarmanRecover._create_recovery_operation") + self.assertIsInstance(call_args[4], KeyError) + self.assertEqual(repr(call_args[4]), "KeyError('operation_id')") + + @patch("logging.warning") + @patch("time.sleep") + @patch.object(BarmanRecover, "_get_request") + def test__get_recovery_operation_status(self, mock_get_request, mock_sleep, mock_logging): + # well formed response + mock_get_request.return_value = {"status": "some status"} + self.assertEqual(self.br._get_recovery_operation_status("some_id"), "some status") + mock_get_request.assert_called_once_with(f"servers/{BARMAN_SERVER}/operations/some_id") + mock_sleep.assert_not_called() + mock_logging.assert_not_called() + + # malformed response + mock_get_request.return_value = {"statuss": "some status"} + + with self.assertRaises(RetriesExceeded) as exc: + self.br._get_recovery_operation_status("some_id") + + self.assertEqual(str(exc.exception), + "Maximum number of retries exceeded for method BarmanRecover._get_recovery_operation_status.") + + self.assertEqual(mock_sleep.call_count, self.br.max_retries) + mock_sleep.assert_has_calls([call(self.br.retry_wait)] * self.br.max_retries) + + self.assertEqual(mock_logging.call_count, self.br.max_retries) + for i in range(mock_logging.call_count): + call_args = mock_logging.mock_calls[i].args + self.assertEqual(len(call_args), 5) + self.assertEqual(call_args[0], "Attempt %d of %d on method %s failed with %r.") + self.assertEqual(call_args[1], i + 1) + self.assertEqual(call_args[2], self.br.max_retries) + self.assertEqual(call_args[3], "BarmanRecover._get_recovery_operation_status") + self.assertIsInstance(call_args[4], KeyError) + self.assertEqual(repr(call_args[4]), "KeyError('status')") + + @patch.object(BarmanRecover, "_get_recovery_operation_status") + @patch("time.sleep") + @patch("logging.info") + @patch("logging.critical") + @patch.object(BarmanRecover, "_create_recovery_operation") + def test_restore_backup(self, mock_create_op, mock_log_critical, mock_log_info, mock_sleep, mock_get_status): + # successful fast restore + mock_create_op.return_value = "some_id" + mock_get_status.return_value = "DONE" + + self.assertTrue(self.br.restore_backup()) + + mock_create_op.assert_called_once() + mock_get_status.assert_called_once_with("some_id") + mock_log_info.assert_called_once_with("Created the recovery operation with ID %s", "some_id") + mock_log_critical.assert_not_called() + mock_sleep.assert_not_called() + + # successful slow restore + mock_create_op.reset_mock() + mock_get_status.reset_mock() + mock_log_info.reset_mock() + mock_get_status.side_effect = ["IN_PROGRESS"] * 20 + ["DONE"] + + self.assertTrue(self.br.restore_backup()) + + mock_create_op.assert_called_once() + + self.assertEqual(mock_get_status.call_count, 21) + mock_get_status.assert_has_calls([call("some_id")] * 21) + + self.assertEqual(mock_log_info.call_count, 21) + mock_log_info.assert_has_calls([call("Created the recovery operation with ID %s", "some_id")] + + [call("Recovery operation %s is still in progress", "some_id")] * 20) + + mock_log_critical.assert_not_called() + + self.assertEqual(mock_sleep.call_count, 20) + mock_sleep.assert_has_calls([call(LOOP_WAIT)] * 20) + + # failed fast restore + mock_create_op.reset_mock() + mock_get_status.reset_mock() + mock_log_info.reset_mock() + mock_sleep.reset_mock() + mock_get_status.side_effect = None + mock_get_status.return_value = "FAILED" + + self.assertFalse(self.br.restore_backup()) + + mock_create_op.assert_called_once() + mock_get_status.assert_called_once_with("some_id") + mock_log_info.assert_called_once_with("Created the recovery operation with ID %s", "some_id") + mock_log_critical.assert_not_called() + mock_sleep.assert_not_called() + + # failed slow restore + mock_create_op.reset_mock() + mock_get_status.reset_mock() + mock_log_info.reset_mock() + mock_sleep.reset_mock() + mock_get_status.side_effect = ["IN_PROGRESS"] * 20 + ["FAILED"] + + self.assertFalse(self.br.restore_backup()) + + mock_create_op.assert_called_once() + + self.assertEqual(mock_get_status.call_count, 21) + mock_get_status.assert_has_calls([call("some_id")] * 21) + + self.assertEqual(mock_log_info.call_count, 21) + mock_log_info.assert_has_calls([call("Created the recovery operation with ID %s", "some_id")] + + [call("Recovery operation %s is still in progress", "some_id")] * 20) + + mock_log_critical.assert_not_called() + + self.assertEqual(mock_sleep.call_count, 20) + mock_sleep.assert_has_calls([call(LOOP_WAIT)] * 20) + + # create retries exceeded + mock_log_info.reset_mock() + mock_sleep.reset_mock() + mock_create_op.side_effect = RetriesExceeded + mock_get_status.side_effect = None + + with self.assertRaises(SystemExit) as exc: + self.assertIsNone(self.br.restore_backup()) + + self.assertEqual(exc.exception.code, ExitCode.HTTP_RESPONSE_MALFORMED) + mock_log_info.assert_not_called() + mock_log_critical.assert_called_once_with("Maximum number of retries exceeded, exiting.") + mock_sleep.assert_not_called() + + # get status retries exceeded + mock_create_op.reset_mock() + mock_create_op.side_effect = None + mock_log_critical.reset_mock() + mock_log_info.reset_mock() + mock_get_status.side_effect = RetriesExceeded + + with self.assertRaises(SystemExit) as exc: + self.assertIsNone(self.br.restore_backup()) + + self.assertEqual(exc.exception.code, ExitCode.HTTP_RESPONSE_MALFORMED) + mock_log_info.assert_called_once_with("Created the recovery operation with ID %s", "some_id") + mock_log_critical.assert_called_once_with("Maximum number of retries exceeded, exiting.") + mock_sleep.assert_not_called() + + +class TestMain(unittest.TestCase): + + @patch("logging.basicConfig") + def test_set_up_logging(self, mock_log_config): + log_file = "/path/to/some/file.log" + set_up_logging(log_file) + mock_log_config.assert_called_once_with(filename=log_file, level=logging.INFO, + format="%(asctime)s %(levelname)s: %(message)s") + + @patch("logging.critical") + @patch("logging.info") + @patch("patroni.scripts.barman_recover.set_up_logging") + @patch("patroni.scripts.barman_recover.BarmanRecover") + @patch("patroni.scripts.barman_recover.ArgumentParser") + def test_main(self, mock_arg_parse, mock_br, mock_set_up_log, mock_log_info, mock_log_critical): + # successful restore + args = MagicMock() + mock_arg_parse.return_value.parse_known_args.return_value = (args, None) + mock_br.return_value.restore_backup.return_value = True + + with self.assertRaises(SystemExit) as exc: + main() + + mock_arg_parse.assert_called_once() + mock_set_up_log.assert_called_once_with(args.log_file) + mock_br.assert_called_once_with(args.api_url, args.barman_server, args.backup_id, args.ssh_command, + args.data_directory, args.loop_wait, args.retry_wait, args.max_retries, + args.cert_file, args.key_file) + mock_log_info.assert_called_once_with("Recovery operation finished successfully.") + mock_log_critical.assert_not_called() + self.assertEqual(exc.exception.code, ExitCode.RECOVERY_DONE) + + # failed restore + mock_arg_parse.reset_mock() + mock_set_up_log.reset_mock() + mock_br.reset_mock() + mock_log_info.reset_mock() + mock_br.return_value.restore_backup.return_value = False + + with self.assertRaises(SystemExit) as exc: + main() + + mock_arg_parse.assert_called_once() + mock_set_up_log.assert_called_once_with(args.log_file) + mock_br.assert_called_once_with(args.api_url, args.barman_server, args.backup_id, args.ssh_command, + args.data_directory, args.loop_wait, args.retry_wait, args.max_retries, + args.cert_file, args.key_file) + mock_log_info.assert_not_called() + mock_log_critical.assert_called_once_with("Recovery operation failed.") + self.assertEqual(exc.exception.code, ExitCode.RECOVERY_FAILED)