diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 86982ba..c8444b5 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -3,6 +3,13 @@ Changelog ========= +1.8.3 (2026-06-18) +================== + +- Broker graceful shutdown with task tracking +- Broker fixes and updates +- Signal and queue hash over name + 1.8.2 (2026-06-11) ================== diff --git a/microagent/__init__.py b/microagent/__init__.py index 0aed7c7..e405978 100644 --- a/microagent/__init__.py +++ b/microagent/__init__.py @@ -1,4 +1,4 @@ -__version__ = '1.8.2' +__version__ = '1.8.3' import importlib import json diff --git a/microagent/abc.py b/microagent/abc.py index 37e047d..688ce6f 100644 --- a/microagent/abc.py +++ b/microagent/abc.py @@ -39,5 +39,8 @@ class BrokerProtocol(Protocol): uid: str log: logging.Logger + async def close(self) -> None: + ... + def __getattr__(self, name: str) -> QueueProtocol: ... diff --git a/microagent/agent.py b/microagent/agent.py index 65f0d2a..718c240 100644 --- a/microagent/agent.py +++ b/microagent/agent.py @@ -210,6 +210,8 @@ async def stop(self) -> None: cron_task.cancel() if self.bus: await self.bus.close() + if self.broker: + await self.broker.close() await self.hook.pre_stop() @staticmethod diff --git a/microagent/broker.py b/microagent/broker.py index 2722e9c..f8d8c1e 100644 --- a/microagent/broker.py +++ b/microagent/broker.py @@ -56,6 +56,7 @@ async def example_read_queue(self, **kwargs): email_agent = EmailAgent(broker=broker) await email_agent.start() ''' +import asyncio import logging import uuid from abc import abstractmethod @@ -108,6 +109,17 @@ class AbstractQueueBroker(BrokerProtocol): log: logging.Logger = logging.getLogger('microagent.broker') _bindings: dict[str, Consumer] = field(default_factory=dict) + _background_tasks: set[asyncio.Task] = field(default_factory=set) + + async def close(self) -> None: + ''' + Graceful shutdown. Cancel all background tasks. + ''' + tasks = list(self._background_tasks) + for task in tasks: + task.cancel() + if tasks: + await asyncio.wait(tasks) def __getattr__(self, name: str) -> 'BoundQueue': return BoundQueue(self, Queue.get(name)) diff --git a/microagent/queue.py b/microagent/queue.py index f5e55e0..0ff4d05 100644 --- a/microagent/queue.py +++ b/microagent/queue.py @@ -87,7 +87,7 @@ def __eq__(self, other: object) -> bool: return self.name == other.name def __hash__(self) -> int: - return id(self) + return id(self.name) @classmethod def set_jsonlib(cls, jsonlib: ModuleType) -> None: diff --git a/microagent/signal.py b/microagent/signal.py index 0faf180..608c57d 100644 --- a/microagent/signal.py +++ b/microagent/signal.py @@ -89,7 +89,7 @@ def __eq__(self, other: object) -> bool: return self.name == other.name def __hash__(self) -> int: - return id(self) + return id(self.name) @classmethod def set_jsonlib(cls, jsonlib: ModuleType) -> None: diff --git a/microagent/tools/amqp.py b/microagent/tools/amqp.py index 8471797..7776f7f 100644 --- a/microagent/tools/amqp.py +++ b/microagent/tools/amqp.py @@ -60,15 +60,30 @@ async def example_read_queue(self, amqp, **data): ''' connection: AbstractConnection = field(init=False) sending_channel: AbstractChannel | None = None + _managed_connections: list['ManagedConnection'] = field(default_factory=list, init=False) + _closing: bool = field(default=False, init=False) def __post_init__(self) -> None: self.connection = ReConnection(self.reconnect, self.dsn) async def reconnect(self) -> None: + if self._closing: + return self.connection = ReConnection(self.reconnect, self.dsn) await self.connection.connect() log.info('Reconnect "%s"', self.connection) + async def close(self) -> None: + self._closing = True + for mc in self._managed_connections: + mc.closed = True + await mc.close() + if self.sending_channel and not self.sending_channel.is_closed: + await self.sending_channel.close() + if self.connection.is_opened: + await self.connection.close() + await super().close() + async def get_channel(self) -> AbstractChannel: ''' Takes a channel from the pool or a new one, performs a lazy connection if required. @@ -103,17 +118,17 @@ async def send(self, name: str, message: str, exchange: str = '', topic: str | N async def bind(self, name: str) -> None: consumer = self._bindings[name] - await ManagedConnection( + mc = ManagedConnection( dsn=self.dsn, consumer=consumer, handler=self._amqp_wrapper(consumer) - ).bind() + ) + self._managed_connections.append(mc) + await mc.bind() def _amqp_wrapper(self, consumer: Consumer) -> Callable[[DeliveredMessage], Awaitable[None]]: async def _wrapper(message: DeliveredMessage) -> None: - if not (data := self.prepared_data(consumer, message.body)): - log.debug('Calling %s by %s without data', consumer, consumer.queue.name) - return + data = self.prepared_data(consumer, message.body) log.debug('Calling %s by %s with %s', consumer, consumer.queue.name, str(data).encode('utf-8')) @@ -126,7 +141,7 @@ async def _wrapper(message: DeliveredMessage) -> None: if consumer.options.get('autoack', True) and message.delivery_tag: await message.channel.basic_ack(delivery_tag=message.delivery_tag) - except TypeError: + except TypeError: # when data mismatch handler args (stay data in queue) log.exception('Call %s failed', consumer) except asyncio.TimeoutError: @@ -185,12 +200,15 @@ class ManagedConnection: handler: Callable bind_attempts: int = 0 bind_running: bool = False + closed: bool = False + _connection: AbstractConnection | None = None + _channel: AbstractChannel | None = None async def bind(self) -> None: ''' Start connection and bind consumer ''' - connection = ReConnection(self.rebind, self.dsn) - await connection.connect() - channel = await connection.channel() + self._connection = ReConnection(self.rebind, self.dsn) + await self._connection.connect() + self._channel = await self._connection.channel() queue_name = self.consumer.queue.name exchange_name = self.consumer.queue.exchange @@ -199,27 +217,35 @@ async def bind(self) -> None: log.warning('Declare queue "%s" with exchange_name "%s" topics "%s"', queue_name, exchange_name, topics) - await channel.queue_declare(queue_name) + await self._channel.queue_declare(queue_name) if topics: # topics - await channel.exchange_declare(exchange=exchange_name, exchange_type='topic') + await self._channel.exchange_declare(exchange=exchange_name, exchange_type='topic') for topic in topics: - await channel.queue_bind(queue=queue_name, exchange=exchange_name, routing_key=topic) + await self._channel.queue_bind(queue=queue_name, exchange=exchange_name, routing_key=topic) elif exchange_name: # fanout - await channel.exchange_declare(exchange=exchange_name, exchange_type='fanout') - await channel.queue_bind(queue=queue_name, exchange=exchange_name) + await self._channel.exchange_declare(exchange=exchange_name, exchange_type='fanout') + await self._channel.queue_bind(queue=queue_name, exchange=exchange_name) + + await self._channel.basic_consume(queue_name, self.handler) - await channel.basic_consume(queue_name, self.handler) + async def close(self) -> None: + if self._channel and not self._channel.is_closed: + await self._channel.close() + if self._connection and self._connection.is_opened: + await self._connection.close() async def rebind(self) -> bool: - if self.bind_running: - log.exception('Already rebinding queue "%s"', self.consumer.queue.name) + if self.closed or self.bind_running: + log.warning('Already rebinding queue "%s"', self.consumer.queue.name) return False if self.bind_attempts > REBIND_ATTEMPTS: - log.exception('Failed all attempts to rebind queue "%s"', self.consumer.queue.name) + log.error('Failed all attempts to rebind queue "%s"', self.consumer.queue.name) return False + self.bind_running = True + await asyncio.sleep((self.bind_attempts ** 2) * REBIND_BASE_DELAY) self.bind_attempts += 1 diff --git a/microagent/tools/kafka.py b/microagent/tools/kafka.py index e621c8c..056f1d1 100644 --- a/microagent/tools/kafka.py +++ b/microagent/tools/kafka.py @@ -44,26 +44,27 @@ async def example_read_queue(self, kafka, **data): ''' addr: str = field(init=False) producer: aiokafka.AIOKafkaProducer = field(init=False) + _producer_started: bool = field(default=False, init=False) def __post_init__(self) -> None: self.addr = parse.urlparse(self.dsn).netloc self.producer = aiokafka.AIOKafkaProducer(bootstrap_servers=self.addr) async def send(self, name: str, message: str, **kwargs: Any) -> None: - await self.producer.start() + if not self._producer_started: + await self.producer.start() + self._producer_started = True + await self.producer.send_and_wait(name, bytes(message, 'utf8'), **kwargs) - try: - await self.producer.send_and_wait(name, bytes(message, 'utf8'), **kwargs) - - finally: - await self.producer.stop() - _loop = asyncio.get_running_loop() - self.producer = aiokafka.AIOKafkaProducer(loop=_loop, bootstrap_servers=self.addr) + async def close(self) -> None: + await self.producer.stop() + await super().close() async def bind(self, name: str) -> None: - loop = asyncio.get_running_loop() - kafka_consumer = aiokafka.AIOKafkaConsumer(name, loop=loop, bootstrap_servers=self.addr) - asyncio.create_task(self._kafka_wrapper(kafka_consumer, name)) + kafka_consumer = aiokafka.AIOKafkaConsumer(name, bootstrap_servers=self.addr) + task = asyncio.create_task(self._kafka_wrapper(kafka_consumer, name)) + self._background_tasks.add(task) + task.add_done_callback(self._background_tasks.discard) async def _kafka_wrapper(self, kafka_consumer: aiokafka.AIOKafkaConsumer, name: str) -> None: consumer = self._bindings[name] @@ -73,7 +74,9 @@ async def _kafka_wrapper(self, kafka_consumer: aiokafka.AIOKafkaConsumer, name: async for msg in kafka_consumer: data = self.prepared_data(consumer, msg.value) data['kafka'] = msg - asyncio.create_task(self._handle(consumer, data)) + htask = asyncio.create_task(self._handle(consumer, data)) + self._background_tasks.add(htask) + htask.add_done_callback(self._background_tasks.discard) finally: await kafka_consumer.stop() diff --git a/microagent/tools/mocks.py b/microagent/tools/mocks.py index 75429d8..5d713ec 100644 --- a/microagent/tools/mocks.py +++ b/microagent/tools/mocks.py @@ -68,14 +68,15 @@ def __init__(self, /, *args: Any, **kw: Any) -> None: self._stuff: dict[str, BoundQueueMock] = {} self.bind_consumer = AsyncMock() self.send = AsyncMock() + self.close = AsyncMock() self.dsn = '' self.uid = '' def __str__(self) -> str: - return f'' def __getattr__(self, name: str) -> BoundQueueMock: - if name.startswith('_') or name in {'bind_consumer', 'send'}: + if name.startswith('_') or name in {'bind_consumer', 'send', 'close'}: return super().__getattr__(name) self._stuff[name] = self._stuff.get(name, BoundQueueMock()) return self._stuff[name] diff --git a/microagent/tools/redis.py b/microagent/tools/redis.py index 1271d39..c957c3b 100644 --- a/microagent/tools/redis.py +++ b/microagent/tools/redis.py @@ -2,7 +2,6 @@ :ref:`Signal Bus ` and :ref:`Queue Broker ` based on :redis:`redis <>`. ''' import asyncio -import inspect import time from collections import defaultdict from dataclasses import dataclass, field @@ -110,33 +109,32 @@ def new_connection(self) -> Redis: return Redis.from_url(self.dsn, decode_responses=True) async def send(self, name: str, message: str, **kwargs: Any) -> None: - ret = self.connection.rpush(name, message) - - if inspect.isawaitable(ret): - await ret + await self.connection.rpush(name, message) async def queue_length(self, name: str, **options: Any) -> int: - ret = self.connection.llen(name) - - if inspect.isawaitable(ret): - return await ret - - return ret + return await self.connection.llen(name) async def bind(self, name: str) -> None: - _loop = asyncio.get_running_loop() - _loop.call_later(self.BIND_TIME, lambda: asyncio.create_task(self._wait(name))) + task = asyncio.create_task(self._wait(name)) + self._background_tasks.add(task) + task.add_done_callback(self._background_tasks.discard) async def _wait(self, name: str) -> None: - conn = await self.new_connection() - while True: - if data := await conn.blpop(name, self.WAIT_TIME): - _, data = data - asyncio.create_task(self._handler(name, data)) # type: ignore[arg-type, unused-ignore] + await asyncio.sleep(self.BIND_TIME) + conn = self.new_connection() + try: + while True: + if data := await conn.blpop(name, self.WAIT_TIME): + _, data = data + htask = asyncio.create_task(self._handler(name, data)) + self._background_tasks.add(htask) + htask.add_done_callback(self._background_tasks.discard) + finally: + await conn.aclose() async def rollback(self, name: str, data: str) -> None: - _hash = str(hash(name)) + str(hash(data)) - attempt = self._rollbacks[_hash] + _key = (name, data) + attempt = self._rollbacks[_key] if attempt > self.ROLLBACK_ATTEMPTS: self.log.error('Rollback limit exceeded on queue "%s" with data: %s', name, data) @@ -147,7 +145,7 @@ async def rollback(self, name: str, data: str) -> None: _loop = asyncio.get_running_loop() _loop.call_later(attempt ** 2, lambda: asyncio.create_task(self.send(name, data))) - self._rollbacks[_hash] += 1 + self._rollbacks[_key] += 1 async def _handler(self, name: str, data: str) -> None: consumer = self._bindings[name] @@ -156,9 +154,9 @@ async def _handler(self, name: str, data: str) -> None: try: await asyncio.wait_for(consumer.handler(**_data), consumer.timeout) - except Exception: - self.log.exception('Call %s failed', consumer.queue.name) - await self.rollback(consumer.queue.name, data) except asyncio.TimeoutError: self.log.error('TimeoutError: %s %.2f', consumer, time.monotonic() - timer) await self.rollback(name, data) + except Exception: + self.log.exception('Call %s failed', consumer.queue.name) + await self.rollback(consumer.queue.name, data) diff --git a/tests/test_tools_amqp.py b/tests/test_tools_amqp.py index 6188d27..8ebad32 100644 --- a/tests/test_tools_amqp.py +++ b/tests/test_tools_amqp.py @@ -136,6 +136,9 @@ async def test_broker_rebind_ok_already(monkeypatch, channel): assert not await conn.rebind() + conn.bind_running = False + assert await conn.rebind() + async def test_broker_rebind_ok_attempts(monkeypatch, channel): queue = Queue(name='test_queue') @@ -204,8 +207,8 @@ async def test_broker_wrapper_ok_nodata(channel): await broker._amqp_wrapper(consumer)(message) - channel.basic_ack.assert_not_called() - handler.assert_not_called() + channel.basic_ack.assert_called() + handler.assert_called_once() async def test_broker_wrapper_ok_type_err(channel): diff --git a/tests/test_tools_kafka.py b/tests/test_tools_kafka.py index b085586..48cd8fe 100644 --- a/tests/test_tools_kafka.py +++ b/tests/test_tools_kafka.py @@ -71,10 +71,17 @@ async def test_broker_send_ok(kafka_producer): broker = KafkaBroker('kafka://localhost') await broker.send(queue.name, '{}') - broker.producer.start.assert_called() + broker.producer.start.assert_called_once() broker.producer.send_and_wait.assert_called_once_with(queue.name, b'{}') - broker.producer.stop.assert_called() + broker.producer.stop.assert_not_called() - broker.producer._closed = False await broker.send(queue.name, '{}') - broker.producer.start.assert_called() + assert broker.producer.send_and_wait.call_count == 2 + + +async def test_broker_close_stops_producer_ok(kafka_producer): + broker = KafkaBroker('kafka://localhost') + await broker.send('test_queue', '{}') + + await broker.close() + broker.producer.stop.assert_called_once() diff --git a/tests/test_tools_redis.py b/tests/test_tools_redis.py index e49083c..ba45d03 100644 --- a/tests/test_tools_redis.py +++ b/tests/test_tools_redis.py @@ -142,7 +142,7 @@ async def test_broker_rollback_many_attempts_ok(): queue = Queue(name='test_queue') broker = RedisBroker('redis://localhost') broker.send = AsyncMock() - broker._rollbacks[str(hash(queue.name)) + str(hash('{}'))] = 4 + broker._rollbacks[(queue.name, '{}')] = 4 await broker.rollback(queue.name, '{}') broker.send.assert_not_called() await asyncio.sleep(.01)