Refactor daemon entrypoints (#2697)

- abstract_main only creates Config object using the passed configfile
  and instantiates the passed daemon class
- common args parser is extracted into a separate func that is called
  from daemons' main funcs (specific args can be added afterwards)

Co-authored-by: Alexander Kukushkin <[email protected]>
This commit is contained in:
Polina Bungina
2023-05-26 15:13:04 +02:00
committed by GitHub
co-authored by Alexander Kukushkin
parent 101ea10e98
commit 37fffa618f
4 changed files with 59 additions and 43 deletions
+27 -6
View File
@@ -1,11 +1,13 @@
import logging import logging
import os import os
import signal import signal
import sys
import time import time
from argparse import Namespace
from typing import Any, Dict, Optional, TYPE_CHECKING from typing import Any, Dict, Optional, TYPE_CHECKING
from patroni.daemon import AbstractPatroniDaemon, abstract_main from patroni.daemon import AbstractPatroniDaemon, abstract_main, get_base_arg_parser
if TYPE_CHECKING: # pragma: no cover if TYPE_CHECKING: # pragma: no cover
from .config import Config from .config import Config
@@ -133,21 +135,40 @@ class Patroni(AbstractPatroniDaemon):
logger.exception('Exception during Ha.shutdown') logger.exception('Exception during Ha.shutdown')
def patroni_main() -> None: def patroni_main(configfile: str) -> None:
from multiprocessing import freeze_support from multiprocessing import freeze_support
from patroni.validator import schema
freeze_support() freeze_support()
abstract_main(Patroni, schema) abstract_main(Patroni, configfile)
def process_arguments() -> Namespace:
parser = get_base_arg_parser()
parser.add_argument('--validate-config', action='store_true', help='Run config validator and exit')
args = parser.parse_args()
if args.validate_config:
from patroni.validator import schema
from patroni.config import Config, ConfigParseError
try:
Config(args.configfile, validator=schema)
sys.exit()
except ConfigParseError as e:
sys.exit(e.value)
return args
def main() -> None: def main() -> None:
from patroni import check_psycopg from patroni import check_psycopg
args = process_arguments()
check_psycopg() check_psycopg()
if os.getpid() != 1: if os.getpid() != 1:
return patroni_main() return patroni_main(args.configfile)
# Patroni started with PID=1, it looks like we are in the container # Patroni started with PID=1, it looks like we are in the container
from types import FrameType from types import FrameType
@@ -180,7 +201,7 @@ def main() -> None:
signal.signal(signal.SIGTERM, passtochild) signal.signal(signal.SIGTERM, passtochild)
import multiprocessing import multiprocessing
patroni = multiprocessing.Process(target=patroni_main) patroni = multiprocessing.Process(target=patroni_main, args=(args.configfile,))
patroni.start() patroni.start()
pid = patroni.pid pid = patroni.pid
patroni.join() patroni.join()
+21 -29
View File
@@ -6,6 +6,7 @@ Currently it is only used for the main "Thread" of ``patroni`` and ``patroni_raf
from __future__ import print_function from __future__ import print_function
import abc import abc
import argparse
import os import os
import signal import signal
import sys import sys
@@ -15,7 +16,22 @@ from typing import Any, Optional, Type, TYPE_CHECKING
if TYPE_CHECKING: # pragma: no cover if TYPE_CHECKING: # pragma: no cover
from .config import Config from .config import Config
from .validator import Schema
def get_base_arg_parser() -> argparse.ArgumentParser:
"""Create a basic argument parser with the arguments used for both patroni and raft controller daemon.
:returns: 'argparse.ArgumentParser' object
"""
from .config import Config
from .version import __version__
parser = argparse.ArgumentParser()
parser.add_argument('--version', action='version', version='%(prog)s {0}'.format(__version__))
parser.add_argument('configfile', nargs='?', default='',
help='Patroni may also read the configuration from the {0} environment variable'
.format(Config.PATRONI_CONFIG_VARIABLE))
return parser
class AbstractPatroniDaemon(abc.ABC): class AbstractPatroniDaemon(abc.ABC):
@@ -141,41 +157,17 @@ class AbstractPatroniDaemon(abc.ABC):
self.logger.shutdown() self.logger.shutdown()
def abstract_main(cls: Type[AbstractPatroniDaemon], validator: Optional['Schema'] = None) -> None: def abstract_main(cls: Type[AbstractPatroniDaemon], configfile: str) -> None:
"""Create the main entry point of a given daemon process. """Create the main entry point of a given daemon process.
Expose a basic argument parser, parse the command-line arguments, and run the given daemon process.
:param cls: a class that should inherit from :class:`AbstractPatroniDaemon`. :param cls: a class that should inherit from :class:`AbstractPatroniDaemon`.
:param validator: used to validate the daemon configuration schema, if requested by the user through :param configfile:
``--validate-config`` CLI option.
""" """
import argparse
from .config import Config, ConfigParseError from .config import Config, ConfigParseError
from .version import __version__
parser = argparse.ArgumentParser()
parser.add_argument('--version', action='version', version='%(prog)s {0}'.format(__version__))
if validator:
parser.add_argument('--validate-config', action='store_true', help='Run config validator and exit')
parser.add_argument('configfile', nargs='?', default='',
help='Patroni may also read the configuration from the {0} environment variable'
.format(Config.PATRONI_CONFIG_VARIABLE))
args = parser.parse_args()
validate_config = validator and args.validate_config
try: try:
if validate_config: config = Config(configfile)
Config(args.configfile, validator=validator)
sys.exit()
config = Config(args.configfile)
except ConfigParseError as e: except ConfigParseError as e:
if e.value: sys.exit(e.value)
print(e.value, file=sys.stderr)
if not validate_config:
parser.print_help()
sys.exit(1)
controller = cls(config) controller = cls(config)
try: try:
+5 -2
View File
@@ -1,7 +1,7 @@
import logging import logging
from .config import Config from .config import Config
from .daemon import AbstractPatroniDaemon, abstract_main from .daemon import AbstractPatroniDaemon, abstract_main, get_base_arg_parser
from .dcs.raft import KVStoreTTL from .dcs.raft import KVStoreTTL
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -27,4 +27,7 @@ class RaftController(AbstractPatroniDaemon):
def main() -> None: def main() -> None:
abstract_main(RaftController) parser = get_base_arg_parser()
args = parser.parse_args()
abstract_main(RaftController, args.configfile)
+6 -6
View File
@@ -15,7 +15,7 @@ from patroni.exceptions import DCSError
from patroni.postgresql import Postgresql from patroni.postgresql import Postgresql
from patroni.postgresql.config import ConfigHandler from patroni.postgresql.config import ConfigHandler
from patroni import check_psycopg from patroni import check_psycopg
from patroni.__main__ import Patroni, main as _main, patroni_main from patroni.__main__ import Patroni, main as _main
from threading import Thread from threading import Thread
from . import psycopg_connect, SleepException from . import psycopg_connect, SleepException
@@ -52,14 +52,14 @@ class TestPatroni(unittest.TestCase):
@patch('sys.argv', ['patroni.py']) @patch('sys.argv', ['patroni.py'])
def test_no_config(self): def test_no_config(self):
self.assertRaises(SystemExit, patroni_main) self.assertRaises(SystemExit, _main)
@patch('sys.argv', ['patroni.py', '--validate-config', 'postgres0.yml']) @patch('sys.argv', ['patroni.py', '--validate-config', 'postgres0.yml'])
@patch('socket.socket.connect_ex', Mock(return_value=1)) @patch('socket.socket.connect_ex', Mock(return_value=1))
def test_validate_config(self): def test_validate_config(self):
self.assertRaises(SystemExit, patroni_main) self.assertRaises(SystemExit, _main)
with patch.object(config.Config, '__init__', Mock(return_value=None)): with patch.object(config.Config, '__init__', Mock(return_value=None)):
self.assertRaises(SystemExit, patroni_main) self.assertRaises(SystemExit, _main)
@patch('pkgutil.iter_importers', Mock(return_value=[MockFrozenImporter()])) @patch('pkgutil.iter_importers', Mock(return_value=[MockFrozenImporter()]))
@patch('sys.frozen', Mock(return_value=True), create=True) @patch('sys.frozen', Mock(return_value=True), create=True)
@@ -94,11 +94,11 @@ class TestPatroni(unittest.TestCase):
with patch('subprocess.call', Mock(return_value=1)): with patch('subprocess.call', Mock(return_value=1)):
with patch.object(Patroni, 'run', Mock(side_effect=SleepException)): with patch.object(Patroni, 'run', Mock(side_effect=SleepException)):
os.environ['PATRONI_POSTGRESQL_DATA_DIR'] = 'data/test0' os.environ['PATRONI_POSTGRESQL_DATA_DIR'] = 'data/test0'
self.assertRaises(SleepException, patroni_main) self.assertRaises(SleepException, _main)
with patch.object(Patroni, 'run', Mock(side_effect=KeyboardInterrupt())): with patch.object(Patroni, 'run', Mock(side_effect=KeyboardInterrupt())):
with patch('patroni.ha.Ha.is_paused', Mock(return_value=True)): with patch('patroni.ha.Ha.is_paused', Mock(return_value=True)):
os.environ['PATRONI_POSTGRESQL_DATA_DIR'] = 'data/test0' os.environ['PATRONI_POSTGRESQL_DATA_DIR'] = 'data/test0'
patroni_main() _main()
@patch('os.getpid') @patch('os.getpid')
@patch('multiprocessing.Process') @patch('multiprocessing.Process')