From b0ca57efdc49fd68d17788d280137e4207dd62ed Mon Sep 17 00:00:00 2001 From: woutdenolf Date: Sat, 16 Aug 2025 22:39:08 +0200 Subject: [PATCH] print_banner: print actual port in case the port is dynamic --- flower/app.py | 174 +++++++++++++++++++++++++++---- flower/command.py | 36 +++---- tests/unit/__init__.py | 57 +++++++--- tests/unit/api/test_control.py | 4 +- tests/unit/api/test_tasks.py | 8 +- tests/unit/test_command.py | 44 +++++--- tests/unit/views/test_monitor.py | 26 ++--- tests/unit/views/test_tasks.py | 16 +-- tests/unit/views/test_workers.py | 38 +++---- 9 files changed, 278 insertions(+), 125 deletions(-) diff --git a/flower/app.py b/flower/app.py index 3427e098a..e09f4223d 100644 --- a/flower/app.py +++ b/flower/app.py @@ -1,5 +1,6 @@ import sys import logging +from urllib.parse import quote from concurrent.futures import ThreadPoolExecutor @@ -37,11 +38,14 @@ class Flower(tornado.web.Application): def __init__(self, options=None, capp=None, events=None, io_loop=None, **kwargs): + handlers = default_handlers if options is not None and options.url_prefix: handlers = [rewrite_handler(h, options.url_prefix) for h in handlers] kwargs.update(handlers=handlers) + super().__init__(**kwargs) + self.options = options or default_options self.io_loop = io_loop or ioloop.IOLoop.instance() self.ssl_options = kwargs.get('ssl_options', None) @@ -49,9 +53,6 @@ def __init__(self, options=None, capp=None, events=None, self.capp = capp or celery.Celery() self.capp.loader.import_default_modules() - self.executor = self.pool_executor_cls(max_workers=self.max_workers) - self.io_loop.set_default_executor(self.executor) - self.inspector = Inspector(self.io_loop, self.capp, self.options.inspect_timeout / 1000.0) self.events = events or Events( @@ -63,33 +64,85 @@ def __init__(self, options=None, capp=None, events=None, io_loop=self.io_loop, max_workers_in_memory=self.options.max_workers, max_tasks_in_memory=self.options.max_tasks) - self.started = False - def start(self): + self._http_server = None + self._executor = None + + def _start_executor(self): + if self._executor is None: + logging.debug("Starting executor...") + ctx = self.pool_executor_cls(max_workers=self.max_workers) + self._executor = ctx.__enter__() # pylint: disable=unnecessary-dunder-call + self.io_loop.set_default_executor(self._executor) + + def _stop_executor(self): + if self._executor is not None: + logging.debug("Stop executor...") + self._executor.__exit__(None, None, None) + self._executor = None + + def _start_events(self): self.events.start() + def _stop_events(self): + self.events.stop() + + def _start_http_server(self): + logging.debug("Starting HTTP server...") if not self.options.unix_socket: - self.listen(self.options.port, address=self.options.address, - ssl_options=self.ssl_options, - xheaders=self.options.xheaders) + http_server = self.listen( + self.options.port, + address=self.options.address, + ssl_options=self.ssl_options, + xheaders=self.options.xheaders + ) else: from tornado.netutil import bind_unix_socket - server = HTTPServer(self) - socket = bind_unix_socket(self.options.unix_socket, mode=0o777) - server.add_socket(socket) - self.started = True - self.update_workers() + http_server = HTTPServer(self) + socket = bind_unix_socket(self.options.unix_socket, mode=0o777) + http_server.add_socket(socket) + self._http_server = http_server + + def _stop_http_server(self): + logging.debug("Stopping HTTP server...") + self.io_loop.run_sync( + self._http_server.close_all_connections, timeout=5 + ) + self._http_server.stop() + self._http_server = None + + def start_server(self): + if self._http_server is not None: + logging.debug("Flower server already started.") + return + logging.debug("Starting Flower server...") + self._start_executor() + self._start_events() + self._start_http_server() + logging.debug("Flower server started.") + + def stop_server(self): + if self._http_server is None: + logging.debug("Flower server already stopped.") + return + logging.debug("Stopping Flower server...") + self._stop_events() + self._stop_http_server() + self._stop_executor() + logging.debug("Flower server stopped.") + + def serve_forever(self): + if not self._http_server: + raise RuntimeError("The server is not running") + logging.debug("Starting event loop...") self.io_loop.start() - def stop(self): - if self.started: - self.events.stop() - logging.debug("Stopping executors...") - self.executor.shutdown(wait=False) - logging.debug("Stopping event loop...") - self.io_loop.stop() - self.started = False + def shutdown(self): + if self._http_server: + raise RuntimeError("The server is still running") + logging.debug("Stopping event loop...") + self.io_loop.stop() @property def transport(self): @@ -101,3 +154,82 @@ def workers(self): def update_workers(self, workername=None): return self.inspector.inspect(workername) + + def _get_scheme(self): + if self.options.unix_socket: + return "http+unix" + if self.ssl_options: + return "https" + return "http" + + def _get_socket(self): + sockets = getattr(self._http_server, "_sockets", None) # pylint: disable=protected-access + if sockets: + return list(sockets.values())[0] + return None + + def _get_domain(self): + if self.options.unix_socket: + raise RuntimeError("UNIX socket") + + sock = self._get_socket() + if sock is not None: + return sock.getsockname()[0] + + return self.options.address or "0.0.0.0" + + def _get_port(self): + if self.options.unix_socket: + raise RuntimeError("UNIX socket") + + sock = self._get_socket() + if sock is not None: + return sock.getsockname()[1] + + return self.options.port + + def _get_authority(self): + if self.options.unix_socket: + return quote(self.options.unix_socket) + + return f"{self._get_domain()}:{self._get_port()}" + + def _get_url_path(self, path=None): + path = path or "" + if not self.options.url_prefix: + return path + + prefix = self.options.url_prefix.strip("/") + return f"/{prefix}{path}" + + def get_url(self, path=None): + path = self._get_url_path(path) + return f"{self._get_scheme()}://{self._get_authority()}{path}" + + # + # For backward compatibility + # + + def start(self): + self.start_server() + self.update_workers() + self.serve_forever() + + def stop(self): + self.stop_server() + self.shutdown() + + @property + def started(self): + return self._http_server is not None + + @started.setter + def started(self, value): + if value: + self.start_server() + else: + self.stop_server() + + @property + def executor(self): + return self._executor diff --git a/flower/command.py b/flower/command.py index 94ed6c7b6..ddb152f08 100644 --- a/flower/command.py +++ b/flower/command.py @@ -50,11 +50,17 @@ def flower(ctx, tornado_argv): atexit.register(flower_app.stop) signal.signal(signal.SIGTERM, sigterm_handler) - if not ctx.obj.quiet: - print_banner(app, 'ssl_options' in settings) + try: + flower_app.start_server() + finally: + # Print the banner even when server failed to start + if not ctx.obj.quiet: + print_banner(flower_app) + + flower_app.update_workers() try: - flower_app.start() + flower_app.serve_forever() except (KeyboardInterrupt, SystemExit): pass @@ -158,24 +164,18 @@ def is_flower_envvar(name): name[len(ENV_VAR_PREFIX):].lower() in default_options -def print_banner(app, ssl): - if not options.unix_socket: - if options.url_prefix: - prefix_str = f'/{options.url_prefix}/' - else: - prefix_str = '' - - logger.info( - "Visit me at http%s://%s:%s%s", 's' if ssl else '', - options.address or '0.0.0.0', options.port, - prefix_str - ) +def print_banner(flower_app): + if not flower_app.options.unix_socket: + url = flower_app.get_url() + logger.info("Visit me at %s", url) else: - logger.info("Visit me via unix socket file: %s", options.unix_socket) + unix_socket = flower_app.options.unix_socket + logger.info("Visit me via unix socket file: %s", unix_socket) - logger.info('Broker: %s', app.connection().as_uri()) + capp = flower_app.capp + logger.info('Broker: %s', capp.connection().as_uri()) logger.info( 'Registered tasks: \n%s', - pformat(sorted(app.tasks.keys())) + pformat(sorted(capp.tasks.keys())) ) logger.debug('Settings: %s', pformat(settings)) diff --git a/tests/unit/__init__.py b/tests/unit/__init__.py index 7412df2b2..dace481a3 100644 --- a/tests/unit/__init__.py +++ b/tests/unit/__init__.py @@ -3,27 +3,60 @@ import celery import tornado.testing -from tornado.ioloop import IOLoop from tornado.options import options +from tornado.httpclient import AsyncHTTPClient, HTTPResponse from flower import command # noqa: F401 side effect - define options from flower.app import Flower -from flower.events import Events from flower.urls import handlers, settings -class AsyncHTTPTestCase(tornado.testing.AsyncHTTPTestCase): +class AsyncHTTPTestCase(tornado.testing.AsyncTestCase): - def _get_celery_app(self): - return celery.Celery() + def setUp(self) -> None: + super().setUp() + self._http_client = AsyncHTTPClient() + self._capp = celery.Celery() + self._start_flower() - def get_app(self, capp=None): - if not capp: - capp = self._get_celery_app() - events = Events(capp, IOLoop.current()) - app = Flower(capp=capp, events=events, - options=options, handlers=handlers, **settings) - return app + def _start_flower(self): + self._app = Flower( + capp=self._capp, + io_loop=self.io_loop, + options=options, + handlers=handlers, + **settings + ) + self._app.start_server() + + def _stop_flower(self): + self._app.stop_server() + + def _restart_flower(self, reset_celery_app=False): + self._stop_flower() + if reset_celery_app: + self._capp = celery.Celery() + self._start_flower() + + def tearDown(self) -> None: + self._http_client.close() + self._app.stop_server() + del self._http_client + del self._app + super().tearDown() + + def fetch( + self, path: str, raise_error: bool = False, **kwargs + ) -> HTTPResponse: + url = self._app.get_url(path) + + def fetch(): + return self._http_client.fetch(url, raise_error=raise_error, **kwargs) + + return self.io_loop.run_sync( + fetch, + timeout=tornado.testing.get_async_test_timeout(), + ) def get(self, url, **kwargs): return self.fetch(url, **kwargs) diff --git a/tests/unit/api/test_control.py b/tests/unit/api/test_control.py index 5c3198c5e..d3155d2f9 100644 --- a/tests/unit/api/test_control.py +++ b/tests/unit/api/test_control.py @@ -16,12 +16,12 @@ def test_unknown_worker(self): class WorkerControlTests(BaseApiTestCase): def setUp(self): - BaseApiTestCase.setUp(self) + super().setUp() self.is_worker = ControlHandler.is_worker ControlHandler.is_worker = lambda *args: True def tearDown(self): - BaseApiTestCase.tearDown(self) + super().tearDown() ControlHandler.is_worker = self.is_worker def test_shutdown(self): diff --git a/tests/unit/api/test_tasks.py b/tests/unit/api/test_tasks.py index 551957d7e..d1b84c090 100644 --- a/tests/unit/api/test_tasks.py +++ b/tests/unit/api/test_tasks.py @@ -97,12 +97,6 @@ def get_task_by_id(events, task_id): class TaskTests(BaseApiTestCase): - def setUp(self): - self.app = super().get_app() - super().setUp() - - def get_app(self, capp=None): - return self.app @patch('flower.api.tasks.tasks', new=MockTasks) def test_task_info(self): @@ -127,7 +121,7 @@ def test_tasks_pagination(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state # Test limit 4 and offset 0 params = dict(limit=4, offset=0, sort_by='name') diff --git a/tests/unit/test_command.py b/tests/unit/test_command.py index c0f7414ef..0a9f30e85 100644 --- a/tests/unit/test_command.py +++ b/tests/unit/test_command.py @@ -74,38 +74,56 @@ def test_autodiscovery(self): - run app.autodiscover_tasks() - create flower command """ - celery_app = self._get_celery_app() + celery_app = self._app.capp with patch.object(celery_app, '_autodiscover_tasks') as autodiscover: celery_app.autodiscover_tasks() - self.get_app(capp=celery_app) + self._restart_flower() self.assertTrue(autodiscover.called) class TestPrintBanner(AsyncHTTPTestCase): + def test_print_banner(self): - celery_app = celery.Celery() with self.assertLogs('', level='INFO') as cm: - print_banner(celery_app, False) + print_banner(self._app) - self.assertTrue('INFO:flower.command:Visit me at http://0.0.0.0:5555' in cm.output) - self.assertTrue('INFO:flower.command:Broker: amqp://guest:**@localhost:5672//' in cm.output) + self.assertIn('INFO:flower.command:Visit me at http://0.0.0.0:5555', cm.output) + self.assertIn('INFO:flower.command:Broker: amqp://guest:**@localhost:5672//', cm.output) def test_print_banner_with_ssl(self): - celery_app = celery.Celery() with self.assertLogs('', level='INFO') as cm: - print_banner(celery_app, True) + self._app.ssl_options = dict(certfile="", keyfile="") + print_banner(self._app) - self.assertTrue('INFO:flower.command:Visit me at https://0.0.0.0:5555' in cm.output) - self.assertTrue('INFO:flower.command:Broker: amqp://guest:**@localhost:5672//' in cm.output) + self.assertIn('INFO:flower.command:Visit me at https://0.0.0.0:5555', cm.output) + self.assertIn('INFO:flower.command:Broker: amqp://guest:**@localhost:5672//', cm.output) def test_print_banner_unix_socket(self): - celery_app = celery.Celery() with self.assertLogs('', level='INFO') as cm, self.mock_option('unix_socket', 'foo'): - print_banner(celery_app, True) + print_banner(self._app) + + self.assertIn('INFO:flower.command:Visit me via unix socket file: foo', cm.output) - self.assertTrue('INFO:flower.command:Visit me via unix socket file: foo' in cm.output) + def test_print_banner_with_dynamic_port(self): + with self.assertLogs('', level='INFO') as cm: + with self.mock_option("port", 0): + max_attempts = 10 + for _ in range(max_attempts): + self._restart_flower() + + port = self._app._get_port() + assert port + if port != 5555: + break + else: + self.fail(f"Port was 5555 after {max_attempts} attempts") + + print_banner(self._app) + + self.assertIn(f'INFO:flower.command:Visit me at http://0.0.0.0:{port}', cm.output) + self.assertIn('INFO:flower.command:Broker: amqp://guest:**@localhost:5672//', cm.output) class TestWarnAboutCeleryArgsUsedInFlowerCommand(AsyncHTTPTestCase): diff --git a/tests/unit/views/test_monitor.py b/tests/unit/views/test_monitor.py index b7d00bd20..b00f44f09 100644 --- a/tests/unit/views/test_monitor.py +++ b/tests/unit/views/test_monitor.py @@ -11,12 +11,6 @@ class PrometheusTests(AsyncHTTPTestCase): - def setUp(self): - self.app = super().get_app() - super().setUp() - - def get_app(self, capp=None): - return self.app def test_metrics(self): state = EventsState() @@ -32,7 +26,7 @@ def test_metrics(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state metrics = self.get('/metrics').body.decode('utf-8') events = dict(re.findall('flower_events_total{task="task1",type="(task-.*)",worker="worker1"} (.*)', metrics)) @@ -61,7 +55,7 @@ def test_task_prefetch_time_metric(self): if e['type'] == 'task-started': e['timestamp'] = task_started state.event(e) - self.app.events.state = state + self._app.events.state = state metrics = self.get('/metrics').body.decode('utf-8') @@ -86,7 +80,7 @@ def test_task_prefetch_time_metric_successful_task_resets_metric_to_zero(self): if e['type'] == 'task-started': e['timestamp'] = task_started state.event(e) - self.app.events.state = state + self._app.events.state = state metrics = self.get('/metrics').body.decode('utf-8') @@ -111,7 +105,7 @@ def test_task_prefetch_time_metric_failed_task_resets_metric_to_zero(self): if e['type'] == 'task-started': e['timestamp'] = task_started state.event(e) - self.app.events.state = state + self._app.events.state = state metrics = self.get('/metrics').body.decode('utf-8') @@ -132,7 +126,7 @@ def test_task_prefetch_time_metric_does_not_compute_prefetch_time_if_task_has_et e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state metrics = self.get('/metrics').body.decode('utf-8') @@ -149,7 +143,7 @@ def test_worker_online_metric_worker_is_offline(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state metrics = self.get('/metrics').body.decode('utf-8') @@ -189,7 +183,7 @@ def test_worker_prefetched_tasks_metric(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state metrics = self.get('/metrics').body.decode('utf-8') @@ -199,12 +193,6 @@ def test_worker_prefetched_tasks_metric(self): class HealthcheckTests(AsyncHTTPTestCase): - def setUp(self): - self.app = super().get_app() - super().setUp() - - def get_app(self, capp=None): - return self.app def test_healthcheck_route(self): response = self.get('/healthcheck').body.decode('utf-8') diff --git a/tests/unit/views/test_tasks.py b/tests/unit/views/test_tasks.py index 9b877d98a..002002a70 100644 --- a/tests/unit/views/test_tasks.py +++ b/tests/unit/views/test_tasks.py @@ -16,12 +16,6 @@ def test_unknown_task(self): class TasksTest(AsyncHTTPTestCase): - def setUp(self): - self.app = super().get_app() - super().setUp() - - def get_app(self, capp=None): - return self.app def test_no_task(self): r = self.get('/tasks') @@ -39,7 +33,7 @@ def test_succeeded_task(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state params = dict(draw=1, start=0, length=10) params['search[value]'] = '' @@ -71,7 +65,7 @@ def test_failed_task(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state params = dict(draw=1, start=0, length=10) params['search[value]'] = '' @@ -109,7 +103,7 @@ def test_sort_runtime(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state params = dict(draw=1, start=0, length=10) params['search[value]'] = '' @@ -157,7 +151,7 @@ def test_sort_incomparable(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state params = dict(draw=1, start=0, length=10) params['search[value]'] = '' @@ -198,7 +192,7 @@ def test_pagination(self): e['clock'] = i e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state params = dict(draw=1, start=0, length=10) params['search[value]'] = '' diff --git a/tests/unit/views/test_workers.py b/tests/unit/views/test_workers.py index 07d37c53c..170a5b886 100644 --- a/tests/unit/views/test_workers.py +++ b/tests/unit/views/test_workers.py @@ -13,12 +13,6 @@ class WorkersTests(AsyncHTTPTestCase): - def setUp(self): - self.app = super().get_app() - super().setUp() - - def get_app(self, capp=None): - return self.app def test_default_page(self): r1 = self.get('/') @@ -45,7 +39,7 @@ def test_single_workers_offline(self): local_received=time.time())) state.event(Event('worker-offline', hostname='worker1', local_received=time.time())) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') table = HtmlTableParser() @@ -65,7 +59,7 @@ def test_purge_offline_workers(self): local_received=time.time())) state.event(Event('worker-offline', hostname='worker1', local_received=time.time())) - self.app.events.state = state + self._app.events.state = state with patch('flower.views.workers.options') as mock_options: mock_options.purge_offline_workers = 0 @@ -82,7 +76,7 @@ def test_single_workers_online(self): state.get_or_create_worker('worker1') state.event(Event('worker-online', hostname='worker1', local_received=time.time())) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') @@ -110,7 +104,7 @@ def test_task_received(self): e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') @@ -140,7 +134,7 @@ def test_task_started(self): e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') @@ -172,7 +166,7 @@ def test_task_succeeded(self): e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') @@ -204,7 +198,7 @@ def test_task_failed(self): e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') @@ -238,7 +232,7 @@ def test_task_retried(self): e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') @@ -271,7 +265,7 @@ def test_tasks(self): e['local_received'] = time.time() state.event(e) - self.app.events.state = state + self._app.events.state = state r = self.get('/workers') @@ -293,7 +287,7 @@ def test_workers_view_json(self): state.get_or_create_worker('worker1') state.event(Event('worker-online', hostname='worker1', local_received=time.time())) - self.app.events.state = state + self._app.events.state = state res = self.get('/workers?json=1') self.assertEqual(200, res.code) @@ -305,9 +299,9 @@ def test_workers_view_refresh(self): state.get_or_create_worker('worker1') state.event(Event('worker-online', hostname='worker1', local_received=time.time())) - self.app.events.state = state + self._app.events.state = state - with patch.object(self.get_app(), "update_workers") as update_workers_mock: + with patch.object(self._app, "update_workers") as update_workers_mock: res = self.get('/workers?refresh=1') self.assertEqual(200, res.code) update_workers_mock.assert_called() @@ -317,17 +311,17 @@ def test_workers_page(self): state.get_or_create_worker('worker1') state.event(Event('worker-online', hostname='worker1', local_received=time.time())) - self.app.events.state = state - self.app.inspector.workers['worker1'] = {'registeres': [], 'active_queues': [], + self._app.events.state = state + self._app.inspector.workers['worker1'] = {'registeres': [], 'active_queues': [], 'stats': {'total': {'tasks.add': 10, 'tasks.sleep': 1, 'tasks.error': 1}, 'broker': {'hostname': 'redis', 'userid': None, 'virtual_host': '/', 'port': 6379}}} - with patch.object(self.get_app(), "update_workers") as update_workers_mock: + with patch.object(self._app, "update_workers") as update_workers_mock: res = self.get('/worker/worker1') self.assertEqual(200, res.code) update_workers_mock.assert_called_once_with(workername='worker1') - with patch.object(self.get_app(), "update_workers") as update_workers_mock: + with patch.object(self._app, "update_workers") as update_workers_mock: res = self.get('/worker/worker2') self.assertEqual(404, res.code) update_workers_mock.assert_called_once_with(workername='worker2')