diff --git a/common/mqueue.py b/common/mqueue.py index c41f6784a..92c0e30b9 100644 --- a/common/mqueue.py +++ b/common/mqueue.py @@ -25,7 +25,7 @@ class Partitioners: fallback_partitioner = DefaultPartitioner() @classmethod - def org_id_partitioner(cls, key: str, all_partitions: [], available_partitions: []): + def org_id_partitioner(cls, key: bytes, all_partitions: [], available_partitions: []): # pylint: disable=broad-except """Kafka producer partitioner, expects org_id as a key and selects partition by modulo.""" org_id_hash = None @@ -143,9 +143,11 @@ async def send_one(self, msg, key=None, headers=None): try: data = bytes(json.dumps(msg).encode("utf-8")) res = await self.client.send_and_wait(self.topic, value=data, key=self._serialize_key(key), headers=headers) - LOGGER.debug(res) + LOGGER.debug("Sent message to Kafka topic %s: %s", self.topic, res) except KafkaError: self.connected = False + LOGGER.exception("Failed to send message to Kafka topic %s", self.topic) + raise async def send_many(self, msg_list, key=None, headers=None): """Send list of messages""" @@ -154,27 +156,31 @@ async def send_many(self, msg_list, key=None, headers=None): for msg in msg_list: data = bytes(json.dumps(msg).encode("utf-8")) res = await self.client.send_and_wait(self.topic, value=data, key=self._serialize_key(key), headers=headers) - LOGGER.debug(res) + LOGGER.debug("Sent message to Kafka topic %s: %s", self.topic, res) except KafkaError: self.connected = False + LOGGER.exception("Failed to send message batch to Kafka topic %s (batch size %d)", self.topic, len(msg_list)) + raise async def send_raw(self, msg: bytes, key=None, headers=None): """Logic around sending raw message""" await self.start() try: res = await self.client.send_and_wait(self.topic, value=msg, key=self._serialize_key(key), headers=headers) - LOGGER.debug(res) + LOGGER.debug("Sent raw message to Kafka topic %s: %s", self.topic, res) except KafkaError: self.connected = False + LOGGER.exception("Failed to send raw message to Kafka topic %s", self.topic) + raise - def send(self, msg, key=None, loop=None, headers=None): + async def send(self, msg, key=None, headers=None): """Sends a message""" - return asyncio.ensure_future(self.send_one(msg, key=key, headers=headers), loop=loop) + await self.send_one(msg, key=key, headers=headers) - def send_list(self, msgs, key=None, loop=None, headers=None): + async def send_list(self, msgs, key=None, headers=None): """Sends list of messages""" - return asyncio.ensure_future(self.send_many(msgs, key=key, headers=headers), loop=loop) + await self.send_many(msgs, key=key, headers=headers) - def send_bytes(self, msg: bytes, key=None, loop=None, headers=None): + async def send_bytes(self, msg: bytes, key=None, headers=None): """Sends a message, where message is already encoded to bytes""" - return asyncio.ensure_future(self.send_raw(msg, key=key, headers=headers), loop=loop) + await self.send_raw(msg, key=key, headers=headers) diff --git a/common/utils.py b/common/utils.py index f985d0299..6692d8c1d 100644 --- a/common/utils.py +++ b/common/utils.py @@ -12,6 +12,7 @@ import pytz import requests +from aiokafka.errors import KafkaError from dateutil.parser import isoparse from prometheus_client import Counter from psycopg2.extras import Json @@ -107,7 +108,7 @@ def format_datetime(datetime_obj): return str(datetime_obj) if datetime_obj else None -def send_msg_to_payload_tracker(producer, msg_dict, status, status_msg=None, loop=None, service="vulnerability"): +async def send_msg_to_payload_tracker(producer, msg_dict, status, status_msg=None, service="vulnerability"): """prepare and send message to payload-tracker""" request_id = (msg_dict.get("platform_metadata", {}) or {}).get("request_id") if not request_id: @@ -123,11 +124,15 @@ def send_msg_to_payload_tracker(producer, msg_dict, status, status_msg=None, loo } if status_msg: tracking_payload["status_msg"] = status_msg - producer.send(tracking_payload, loop=loop) - LOGGER.debug("Sent message to topic %s: %s", producer.topic, str(tracking_payload)) + LOGGER.debug("Sending message to topic %s: %s", producer.topic, str(tracking_payload)) + try: + await producer.send(tracking_payload) + except KafkaError: + LOGGER.error("Sending Kafka producer.topic %s message failed", producer.topic) + pass # Work should compate and the caller should not treat the flow as failed, best-effort -def send_remediations_update(producer, inventory_id: str, cves: list, loop=None) -> None: +async def send_remediations_update(producer, inventory_id: str, cves: list) -> None: """ Send a message in format of remediations application updates using given Kafka producer @@ -138,7 +143,7 @@ def send_remediations_update(producer, inventory_id: str, cves: list, loop=None) loop (asyncio.event_loop, optional): asyncio event loop to be used by producer """ msg = {"host_id": inventory_id, "issues": ["vulnerabilities:{}".format(cve) for cve in cves]} - producer.send(msg, loop=loop) + await producer.send(msg) def ensure_minimal_schema_version(): @@ -292,7 +297,9 @@ async def wrapper(*args, **kwargs): return decorator -def send_notifications(notif_topic, new_sys_vulns, mit_sys_vulns, unmit_sys_vulns, rh_account_id, org_id, inventory_id, display_name=None): +async def send_notifications( + notif_topic, new_sys_vulns, mit_sys_vulns, unmit_sys_vulns, rh_account_id, org_id, inventory_id, display_name=None +): """Sends kafka message to notificator with system_vulnerabilities""" if not new_sys_vulns and not mit_sys_vulns and not unmit_sys_vulns: return @@ -312,10 +319,10 @@ def send_notifications(notif_topic, new_sys_vulns, mit_sys_vulns, unmit_sys_vuln ], } LOGGER.debug("Sending evaluation result to notificator: %s", msg) - notif_topic.send(msg) + await notif_topic.send(msg) -def send_inventory_views( +async def send_inventory_views( inventory_views_topic, request_id, inventory_id, @@ -357,7 +364,7 @@ def send_inventory_views( ("request_id", bytes(request_id, "ascii")), ] LOGGER.debug("Sending evaluation result to inventory views: %s", msg) - inventory_views_topic.send(msg, headers=headers) + await inventory_views_topic.send(msg, headers=headers) def create_task_and_log(coro, logger, loop): diff --git a/evaluator/processor.py b/evaluator/processor.py index 527e644b4..d2bcad7be 100644 --- a/evaluator/processor.py +++ b/evaluator/processor.py @@ -378,8 +378,8 @@ async def _evaluate_system( cves_with_known_exploits += 1 await self._mark_system_evaluated(total_cves, system_platform, conn) - send_remediations_update(self.remediations_results, inventory_id, fixable_sys_vuln_rows) - send_notifications( + await send_remediations_update(self.remediations_results, inventory_id, fixable_sys_vuln_rows) + await send_notifications( self.evaluator_results, new_system_vulns, [], @@ -389,7 +389,7 @@ async def _evaluate_system( inventory_id, system_platform.display_name, ) - send_inventory_views( + await send_inventory_views( self.inventory_views_results, request_id, inventory_id, @@ -415,12 +415,12 @@ async def evaluate_system( await self._evaluate_system(inventory_id, org_id, request_id, request_timestamp, recalc_event_id=recalc_event_id) except EvaluatorException as ex: LOGGER.error(str(ex)) - send_msg_to_payload_tracker(self.payload_tracker, msg, "error", status_msg="evaluation failed", loop=self.loop) + await send_msg_to_payload_tracker(self.payload_tracker, msg, "error", status_msg="evaluation failed") return except VmaasErrorException as ex: LOGGER.error(str(ex)) VMAAS_ERRORS_SKIP.inc() - send_msg_to_payload_tracker(self.payload_tracker, msg, "error", status_msg="evaluation failed", loop=self.loop) + await send_msg_to_payload_tracker(self.payload_tracker, msg, "error", status_msg="evaluation failed") return - send_msg_to_payload_tracker(self.payload_tracker, msg, "success", status_msg="evaluation succeeded", loop=self.loop) + await send_msg_to_payload_tracker(self.payload_tracker, msg, "success", status_msg="evaluation succeeded") diff --git a/grouper/queue.py b/grouper/queue.py index 16ed60e22..024d6b9aa 100644 --- a/grouper/queue.py +++ b/grouper/queue.py @@ -161,14 +161,10 @@ async def _send_for_evaluation(self, item: QueueItem, org_id: str, inventory_id: if (not item.inventory_changed and not item.advisor_changed) and not CFG.disable_optimisation: UNCHANGED_SYSTEM.inc() LOGGER.info("skipping evaluation, system not changed: %s, org_id: %s", inventory_id, org_id) - send_msg_to_payload_tracker( - self.payload_tracker, msg, "success", status_msg="unchanged system, not sending to evaluator", loop=self.loop - ) + await send_msg_to_payload_tracker(self.payload_tracker, msg, "success", status_msg="unchanged system, not sending to evaluator") return CHANGED_SYSTEM.inc() LOGGER.info("sending upload message to evaluator: %s, org_id: %s", inventory_id, org_id) - send_msg_to_payload_tracker( - self.payload_tracker, msg, "processing", status_msg="changed system, sending to evaluator", loop=self.loop - ) - self.evaluator.send(msg) + await send_msg_to_payload_tracker(self.payload_tracker, msg, "processing", status_msg="changed system, sending to evaluator") + await self.evaluator.send(msg) diff --git a/listener/advisor_processor.py b/listener/advisor_processor.py index 1759f3615..27780c34f 100644 --- a/listener/advisor_processor.py +++ b/listener/advisor_processor.py @@ -168,7 +168,7 @@ def _parse_input_metadata(self, msg: AdvisorMsg) -> (str, str, str): msg.msg["input"]["timestamp"], ) - def _send_for_evaluation( + async def _send_for_evaluation( self, org_id: str, inventory_id: str, request_id: str, reporter: str, timestamp: str, import_status: ImportStatus ): """Send message to evaluator to evaluate""" @@ -185,11 +185,11 @@ def _send_for_evaluation( }, "timestamp": timestamp, } - self.grouper.send(msg, loop=self.loop, key=org_id) + await self.grouper.send(msg, key=org_id) - def _send_to_payload_tracker(self, status: str, msg: AdvisorMsg, message=None): + async def _send_to_payload_tracker(self, status: str, msg: AdvisorMsg, message=None): """Send payload tracker message""" - send_msg_to_payload_tracker(self.payload_tracker, msg.msg["input"], status, status_msg=message, loop=self.loop) + await send_msg_to_payload_tracker(self.payload_tracker, msg.msg["input"], status, status_msg=message) async def _process_upload(self, msg: AdvisorMsg): """Process message from advisor""" @@ -221,8 +221,8 @@ async def _process_upload(self, msg: AdvisorMsg): LOGGER.info( "advisor data inserted, system: %s, org_id: %s, reporter: %s, request_id: %s", inventory_id, org_id, reporter, request_id ) - self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) - self._send_to_payload_tracker("received", msg, message="system received from advisor, sending to grouper") + await self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) + await self._send_to_payload_tracker("received", msg, message="system received from advisor, sending to grouper") async def process_msg(self, msg: AdvisorMsg): """Process single advisor msg""" diff --git a/listener/inventory_processor.py b/listener/inventory_processor.py index a1cd6385d..d96fb08de 100644 --- a/listener/inventory_processor.py +++ b/listener/inventory_processor.py @@ -449,7 +449,7 @@ def _parse_input_metadata(self, msg: InventoryMsg) -> (str, str, str): """Extract reporter, request_id and timestamp from inventory message""" return msg.msg["host"]["reporter"], (msg.msg.get("platform_metadata") or {}).get("request_id", ""), msg.msg["timestamp"] - def _send_for_evaluation( + async def _send_for_evaluation( self, org_id: str, inventory_id: str, request_id: str, reporter: str, timestamp: str, import_status: ImportStatus ): """Send message to evaluator to evaluate""" @@ -466,11 +466,11 @@ def _send_for_evaluation( }, "timestamp": timestamp, } - self.grouper.send(msg, loop=self.loop, key=org_id) + await self.grouper.send(msg, key=org_id) - def _send_to_payload_tracker(self, status: str, msg: InventoryMsg, message=None): + async def _send_to_payload_tracker(self, status: str, msg: InventoryMsg, message=None): """Send payload tracker message""" - send_msg_to_payload_tracker(self.payload_tracker, msg.msg, status, status_msg=message, loop=self.loop) + await send_msg_to_payload_tracker(self.payload_tracker, msg.msg, status, status_msg=message) async def _process_upload(self, msg: InventoryMsg): """Process upload message defined by QueueItem""" @@ -544,8 +544,8 @@ async def _process_upload(self, msg: InventoryMsg): LOGGER.info( "Inventory data processed, system: %s, org_id: %s, reporter: %s, request_id: %s", inventory_id, org_id, reporter, request_id ) - self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) - self._send_to_payload_tracker("received", msg, message="system received from inventory, sending to grouper") + await self._send_for_evaluation(org_id, inventory_id, request_id, reporter, timestamp, import_status) + await self._send_to_payload_tracker("received", msg, message="system received from inventory, sending to grouper") async def _process_delete(self, msg: InventoryMsg): """Process inventory delete message""" diff --git a/listener/listener.py b/listener/listener.py index dab63bdb6..7723260b9 100644 --- a/listener/listener.py +++ b/listener/listener.py @@ -152,9 +152,9 @@ async def init(self): """Async constructor""" await self.inventory_msg_processor.init() - def _send_payload_tracker_error(self, msg: dict, reason: str): + async def _send_payload_tracker_error(self, msg: dict, reason: str): # since messages come separetely, inform separetly user about error in message - send_msg_to_payload_tracker(self.payload_tracker, msg, "error", status_msg=reason) + await send_msg_to_payload_tracker(self.payload_tracker, msg, "error", status_msg=reason) async def _consume_inventory_msg(self, msg: dict) -> InventoryMsgType: """Consumes inventory message""" @@ -163,7 +163,7 @@ async def _consume_inventory_msg(self, msg: dict) -> InventoryMsgType: except InvalidInventoryMsg as exc: LOGGER.error("obtained invalid inventory msg: %s", exc) try: - self._send_payload_tracker_error(msg, "obtained invalid inventory msg") + await self._send_payload_tracker_error(msg, "obtained invalid inventory msg") except KeyError: pass return InventoryMsgType.UNKNOWN @@ -185,7 +185,7 @@ async def _consume_advisor_msg(self, msg: dict): except InvalidAdvisorMsg as exc: LOGGER.error("obtained invalid advisor msg: %s", exc) try: - self._send_payload_tracker_error(msg["input"], "obtained invalid advisor msg") + await self._send_payload_tracker_error(msg["input"], "obtained invalid advisor msg") except KeyError: pass return diff --git a/manager/admin_handler.py b/manager/admin_handler.py index b2d820d64..f4c534099 100644 --- a/manager/admin_handler.py +++ b/manager/admin_handler.py @@ -38,6 +38,14 @@ BATCH_SEMAPHORE = asyncio.BoundedSemaphore(CFG.re_evaluation_kafka_batches) +async def _send_recalc_batch_and_release(msgs): + """Send a recalc batch to Kafka and release the batch semaphore""" + try: + await EVALUATOR_QUEUE.send_list(msgs) + finally: + BATCH_SEMAPHORE.release() + + class TaskomaticRun(PutRequest): """PUT to /v1/taskomatic/run""" @@ -433,8 +441,8 @@ def handle_delete(cls, **kwargs): class RecalcBase: @classmethod - def _create_kafka_msg_task(cls, rows, loop): - msgs = [ + def _create_kafka_msg(cls, rows): + return [ { "type": EvaluatorMessageType.RE_EVALUATE_SYSTEM, "host": {"id": str(inventory_id), "org_id": org_id}, @@ -442,8 +450,6 @@ def _create_kafka_msg_task(cls, rows, loop): } for inventory_id, org_id in rows ] - task = EVALUATOR_QUEUE.send_list(msgs, loop=loop) - return task, len(msgs) class RecalcAccounts(PutRequest, RecalcBase): @@ -472,10 +478,9 @@ def handle_put(cls, **kwargs): if not rows: BATCH_SEMAPHORE.release() break - task, msg_count = cls._create_kafka_msg_task(rows, loop) - total_scheduled += msg_count - task.add_done_callback(lambda x: BATCH_SEMAPHORE.release()) - loop.run_until_complete(task) + msgs = cls._create_kafka_msg(rows) + loop.run_until_complete(_send_recalc_batch_and_release(msgs)) + total_scheduled += len(msgs) return f"{total_scheduled} systems scheduled for re-evaluation", 200 @@ -503,8 +508,9 @@ def handle_put(cls, **kwargs): if not rows: return f"{total_scheduled} systems scheduled for re-evaluation", 200 - task, total_scheduled = cls._create_kafka_msg_task(rows, loop) - loop.run_until_complete(task) + msgs = cls._create_kafka_msg(rows) + loop.run_until_complete(EVALUATOR_QUEUE.send_list(msgs)) + total_scheduled = len(msgs) return f"{total_scheduled} systems scheduled for re-evaluation", 200 diff --git a/notificator/notificator_queue.py b/notificator/notificator_queue.py index 79b010708..2e0167a40 100644 --- a/notificator/notificator_queue.py +++ b/notificator/notificator_queue.py @@ -118,7 +118,7 @@ def _create_notif_events(self, cve_id): } ] - def _send_kafka_notif( + async def _send_kafka_notif( self, org_id: str, inventory_id: str, @@ -127,7 +127,6 @@ def _send_kafka_notif( events: list, only_admins=False, ignore_user_preferences=False, - loop=None, ): """ Sends kafka msg to notification kafka. @@ -157,7 +156,7 @@ def _send_kafka_notif( msg["org_id"] = org_id msg["context"]["org_id"] = org_id LOGGER.debug("Sending notification: %s", msg) - self.notifications_topic.send(msg, loop=loop) + await self.notifications_topic.send(msg) async def _register_notified_acc(self, cve_id: int, notif_events: set(NotificationType), rh_account_id): """Registers new notified accounts into notified_accounts table""" @@ -198,7 +197,7 @@ async def _process_normal_queue(self): else: # account-level: existing dedup logic if not await self._is_already_notified(item.rh_account_id, item.cve_id, notif_event): - self._send_kafka_notif(item.org_id, inventory_id, display_name, notif_event.value, cve_events, loop=self.loop) + await self._send_kafka_notif(item.org_id, inventory_id, display_name, notif_event.value, cve_events) LOGGER.info("Sent %s, cve=%s, org_id=%s", notif_event.value, item.cve, item.org_id) SENT_NOTIFICATIONS.inc() new_notified.add(notif_event) @@ -213,7 +212,7 @@ async def _process_normal_queue(self): # send batched system-level notifications (one message per system per notification type) for (inventory_id, org_id, display_name, notif_event), events in system_notif_batch.items(): - self._send_kafka_notif(org_id, inventory_id, display_name, notif_event.value, events, loop=self.loop) + await self._send_kafka_notif(org_id, inventory_id, display_name, notif_event.value, events) LOGGER.info("Sent %s, inventory_id=%s, cve_count=%s", notif_event.value, inventory_id, len(events)) SENT_NOTIFICATIONS.inc() diff --git a/tests/notificator_tests/test_notificator_queue.py b/tests/notificator_tests/test_notificator_queue.py index 40c6830dc..9edf9e93c 100644 --- a/tests/notificator_tests/test_notificator_queue.py +++ b/tests/notificator_tests/test_notificator_queue.py @@ -54,7 +54,7 @@ async def test_queue_valid(asyncpg_pool): event_loop = asyncio.get_running_loop() msgs = [] - def _send_kafka_notif_mock(self, acc_id, event_type, *_, **__): + async def _send_kafka_notif_mock(self, acc_id, event_type, *_, **__): msgs.append((acc_id, event_type)) queue = await TestNotificatorQueue._build_queue(asyncpg_pool, event_loop) @@ -89,7 +89,7 @@ async def test_queue_unknown_cve(asyncpg_pool): event_loop = asyncio.get_running_loop() msgs = [] - def _send_kafka_notif_mock(org_id, inventory_id, display_name, event_type, events, *_, **__): + async def _send_kafka_notif_mock(org_id, inventory_id, display_name, event_type, events, *_, **__): msgs.append((inventory_id, event_type, len(events))) queue = await TestNotificatorQueue._build_queue(asyncpg_pool, event_loop) @@ -126,7 +126,7 @@ async def test_queue_system_notif_batching(asyncpg_pool): event_loop = asyncio.get_running_loop() msgs = [] - def _send_kafka_notif_mock(org_id, inventory_id, display_name, event_type, events, *_, **__): + async def _send_kafka_notif_mock(org_id, inventory_id, display_name, event_type, events, *_, **__): msgs.append((inventory_id, event_type, len(events))) queue = await TestNotificatorQueue._build_queue(asyncpg_pool, event_loop) @@ -177,7 +177,7 @@ async def test_queue_system_notif_batching_same_type(asyncpg_pool): event_loop = asyncio.get_running_loop() msgs = [] - def _send_kafka_notif_mock(org_id, inventory_id, display_name, event_type, events, *_, **__): + async def _send_kafka_notif_mock(org_id, inventory_id, display_name, event_type, events, *_, **__): msgs.append((inventory_id, event_type, len(events))) queue = await TestNotificatorQueue._build_queue(asyncpg_pool, event_loop) diff --git a/tests/vmaas_sync_tests/test_vmaas_sync.py b/tests/vmaas_sync_tests/test_vmaas_sync.py index 61f08d9d3..0045c5d52 100644 --- a/tests/vmaas_sync_tests/test_vmaas_sync.py +++ b/tests/vmaas_sync_tests/test_vmaas_sync.py @@ -14,6 +14,7 @@ from vmaas_sync.vmaas_sync import LOGGER from vmaas_sync.vmaas_sync import _get_last_repobased_eval_tms from vmaas_sync.vmaas_sync import _select_repo_based_inventory_ids +from vmaas_sync.vmaas_sync import _send_recalc_batch_and_release from vmaas_sync.vmaas_sync import _set_last_repobased_eval_tms from vmaas_sync.vmaas_sync import re_evaluate_systems from vmaas_sync.vmaas_sync import sync_cve_md @@ -34,11 +35,11 @@ def send(self, msg, loop=None): LOGGER.info(msg) return asyncio.ensure_future(self._do_nothing(), loop=loop) - def send_list(self, msg_list, loop=None): + async def send_list(self, msg_list, loop=None): # pylint: disable=unused-argument """send from re_evaluate_systems""" for msg in msg_list: LOGGER.info(msg["host"]["id"]) - return asyncio.ensure_future(self._do_nothing(), loop=loop) + await self._do_nothing() class TestVmaasSync: @@ -105,6 +106,39 @@ def test_re_evaluate_repo_based(self, pg_db_conn, monkeypatch, caplog): # pylin self.check_re_eval_repo_based_logs(caplog.records, 0) caplog.clear() + def test_recalc_batches_schedule_with_create_task(self, monkeypatch): + """ + More recalc batches than semaphore slots must finish without hanging + """ + monkeypatch.setattr(vmaas_sync, "BATCH_SEMAPHORE", asyncio.BoundedSemaphore(2)) + + in_flight = 0 + max_in_flight = 0 + send_calls = 0 + + class TrackingQueue: + async def send_list(self, msg_list): # pylint: disable=unused-argument + nonlocal in_flight, max_in_flight, send_calls + in_flight += 1 + max_in_flight = max(max_in_flight, in_flight) + send_calls += 1 + await asyncio.sleep(0.05) + in_flight -= 1 + + monkeypatch.setattr(vmaas_sync, "EVALUATOR_QUEUE", TrackingQueue()) + + loop = vmaas_sync.EVENT_LOOP + batches = [[{"host": {"id": f"inv-{i}", "org_id": "org-1"}}] for i in range(5)] + send_tasks = [] + for msgs in batches: + loop.run_until_complete(vmaas_sync.BATCH_SEMAPHORE.acquire()) + send_tasks.append(loop.create_task(_send_recalc_batch_and_release(msgs))) + + loop.run_until_complete(asyncio.wait_for(asyncio.gather(*send_tasks), timeout=2.0)) + + assert send_calls == 5 + assert max_in_flight == 2 + def test_select_inv_ids_2(self, pg_db_conn): """Test select inventory_ids repo collection""" with pg_db_conn: diff --git a/vmaas_sync/vmaas_sync.py b/vmaas_sync/vmaas_sync.py index 026631f98..6b6c3e3da 100644 --- a/vmaas_sync/vmaas_sync.py +++ b/vmaas_sync/vmaas_sync.py @@ -38,6 +38,14 @@ BATCH_SEMAPHORE = asyncio.BoundedSemaphore(CFG.re_evaluation_kafka_batches) +async def _send_recalc_batch_and_release(msgs): + """Send a recalc batch to Kafka and release the batch semaphore""" + try: + await EVALUATOR_QUEUE.send_list(msgs) + finally: + BATCH_SEMAPHORE.release() + + @dataclass class OSRelease: """OSRelease object.""" @@ -361,8 +369,9 @@ def re_evaluate_systems(): loop = EVENT_LOOP total_scheduled = 0 - futures = [] + send_tasks = [] while True: + LOGGER.debug("Acquiring the semaphore... (%d)", total_scheduled) loop.run_until_complete(BATCH_SEMAPHORE.acquire()) rows = cur.fetchmany(size=CFG.re_evaluation_kafka_batch_size) if not rows: @@ -378,13 +387,11 @@ def re_evaluate_systems(): for inventory_id, _, org_id in rows ] total_scheduled += len(msgs) - future = EVALUATOR_QUEUE.send_list(msgs, loop=loop) - future.add_done_callback(lambda x: BATCH_SEMAPHORE.release()) - futures.append(future) + send_tasks.append(loop.create_task(_send_recalc_batch_and_release(msgs))) - if futures: - LOGGER.info("Waiting for %s Kafka send operations to complete", len(futures)) - loop.run_until_complete(asyncio.gather(*futures)) + if send_tasks: + LOGGER.info("Waiting for %s Kafka send operations to complete", len(send_tasks)) + loop.run_until_complete(asyncio.gather(*send_tasks)) LOGGER.info("%s systems scheduled for re-evaluation", total_scheduled) conn.commit()