Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions CHANGELOG.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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)
==================

Expand Down
2 changes: 1 addition & 1 deletion microagent/__init__.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
__version__ = '1.8.2'
__version__ = '1.8.3'

import importlib
import json
Expand Down
3 changes: 3 additions & 0 deletions microagent/abc.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,5 +39,8 @@ class BrokerProtocol(Protocol):
uid: str
log: logging.Logger

async def close(self) -> None:
...

def __getattr__(self, name: str) -> QueueProtocol:
...
2 changes: 2 additions & 0 deletions microagent/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
12 changes: 12 additions & 0 deletions microagent/broker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand Down
2 changes: 1 addition & 1 deletion microagent/queue.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
2 changes: 1 addition & 1 deletion microagent/signal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
62 changes: 44 additions & 18 deletions microagent/tools/amqp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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'))
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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

Expand Down
27 changes: 15 additions & 12 deletions microagent/tools/kafka.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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()
Expand Down
5 changes: 3 additions & 2 deletions microagent/tools/mocks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'<BrokerMock id={id(self)}'
return f'<BrokerMock id={id(self)}>'

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]
46 changes: 22 additions & 24 deletions microagent/tools/redis.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@
:ref:`Signal Bus <bus>` and :ref:`Queue Broker <broker>` based on :redis:`redis <>`.
'''
import asyncio
import inspect
import time
from collections import defaultdict
from dataclasses import dataclass, field
Expand Down Expand Up @@ -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)
Expand All @@ -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]
Expand All @@ -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)
7 changes: 5 additions & 2 deletions tests/test_tools_amqp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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')
Expand Down Expand Up @@ -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):
Expand Down
Loading
Loading