diff --git a/bumper/__init__.py b/bumper/__init__.py index 1fb27fc..319de2f 100644 --- a/bumper/__init__.py +++ b/bumper/__init__.py @@ -134,10 +134,7 @@ async def start(): asyncio.create_task(xmpp_server.start_async_server()) # Wait for helperbot to connect first - while mqtt_helperbot.client is None: - await asyncio.sleep(0.1) - - while not mqtt_helperbot.client.session.transitions.state == "connected": + while not mqtt_helperbot.is_connected: await asyncio.sleep(0.1) # Start web servers @@ -170,14 +167,14 @@ async def shutdown(): await conf_server.stop_server() await conf_server_2.stop_server() - if mqtt_server.broker.transitions.state == "started": - await mqtt_server.broker.shutdown() - elif mqtt_server.broker.transitions.state == "starting": - while mqtt_server.broker.transitions.state == "starting": + if mqtt_server.state == "started": + await mqtt_server.shutdown() + elif mqtt_server.state == "starting": + while mqtt_server.state == "starting": await asyncio.sleep(0.1) - if mqtt_server.broker.transitions.state == "started": - await mqtt_server.broker.shutdown() - await mqtt_helperbot.client.disconnect() + if mqtt_server.state == "started": + await mqtt_server.shutdown() + await mqtt_helperbot.disconnect() if xmpp_server.server: if xmpp_server.server._serving: xmpp_server.server.close() diff --git a/bumper/confserver.py b/bumper/confserver.py index c268395..4c7c008 100644 --- a/bumper/confserver.py +++ b/bumper/confserver.py @@ -175,7 +175,7 @@ class ConfServer: bots = bumper.db_get().table("bots").all() clients = bumper.db_get().table("clients").all() - helperbot = bumper.mqtt_helperbot.client.session.transitions.state + helperbot_connected = bumper.mqtt_helperbot.is_connected mqttserver = bumper.mqtt_server.broker xmppserver = bumper.xmpp_server mq_sessions = [] @@ -193,9 +193,9 @@ class ConfServer: all = { "bots": bots, "clients": clients, - "helperbot": [{"state": helperbot}], + "helperbot": {"connected": helperbot_connected}, "mqtt_server": [ - {"state": mqttserver.transitions.state}, + {"state": bumper.mqtt_server.state}, { "sessions": [ {"count": len(mqttserver._sessions)}, @@ -274,32 +274,24 @@ class ConfServer: return await handler(request) async def restart_Helper(self): - - await bumper.mqtt_helperbot.client.disconnect() + await bumper.mqtt_helperbot.disconnect() asyncio.create_task(bumper.mqtt_helperbot.start()) async def restart_MQTT(self): + loop = asyncio.get_event_loop() - if not ( - bumper.mqtt_server.broker.transitions.state == "stopped" - or bumper.mqtt_server.broker.transitions.state == "not_started" - ): + if bumper.mqtt_server.state not in ["stopped", "not_started"]: # close session writers - this was required so bots would reconnect properly after restarting for sess in list(bumper.mqtt_server.broker._sessions): sessobj = bumper.mqtt_server.broker._sessions[sess][1] if sessobj.session.transitions.state == "connected": await sessobj.writer.close() - # await bumper.mqtt_server.broker.shutdown() - aloop = asyncio.get_event_loop() - aloop.call_later( - 0.1, lambda: asyncio.create_task(bumper.mqtt_server.broker.shutdown()) - ) # In .1 seconds shutdown broker + loop.call_later( + 0.1, lambda: asyncio.create_task(bumper.mqtt_server.shutdown()) + ) - aloop = asyncio.get_event_loop() - aloop.call_later( - 1.5, lambda: asyncio.create_task(bumper.mqtt_server.start()) - ) # In 1.5 seconds start broker + loop.call_later(1.5, lambda: asyncio.create_task(bumper.mqtt_server.start())) async def restart_XMPP(self): bumper.xmpp_server.disconnect() diff --git a/bumper/mqttserver.py b/bumper/mqttserver.py index d2943ce..506f18a 100644 --- a/bumper/mqttserver.py +++ b/bumper/mqttserver.py @@ -3,6 +3,7 @@ import asyncio import json import os +from asyncio import Task from typing import Any, Dict, MutableMapping, Optional, Union import amqtt @@ -50,18 +51,15 @@ class MQTTHelperBot: self._client_id = "helperbot@bumper/helperbot" self._timeout = timeout self._client: Optional[MQTTClient] = None + self._new_messages_task: Optional[Task] = None @property - def commands(self) -> MutableMapping[str, CommandDto]: - return self._commands - - @property - def timeout(self) -> float: - return self._timeout - - @property - def client(self) -> MQTTClient: - return self._client + def is_connected(self) -> bool: + """Return True if client is connected successfully.""" + return ( + self._client is not None + and self._client.session.transitions.state == "connected" + ) async def start(self) -> None: try: @@ -77,20 +75,39 @@ class MQTTHelperBot: await self._client.subscribe( [ ("iot/p2p/+/+/+/+/helperbot/bumper/helperbot/+/+/+", QOS_0), - ("iot/p2p/+", QOS_0), - ("iot/atr/+", QOS_0), ] ) + self._new_messages_task = asyncio.create_task( + self._check_for_new_messages() + ) except Exception as e: helperbotlog.exception(f"{e}") raise e + async def _check_for_new_messages(self): + while True: + try: + message: IncomingApplicationMessage = ( + await self._client.deliver_message() + ) + if message is not None: + topic_split = str(message.topic).split("/") + data_decoded = str(message.data.decode("utf-8")) + if topic_split[10] in self._commands: + self._commands[topic_split[10]].add_response(data_decoded) + except asyncio.CancelledError: + pass + except Exception as ex: # pylint: disable=broad-except + helperbotlog.error( + "An exception occurred during handling new messages", exc_info=True + ) + async def _wait_for_resp( self, command_dto: CommandDto, request_id: str ) -> dict[str, Any]: try: payload = await asyncio.wait_for( - command_dto.wait_for_response(), timeout=self.timeout + command_dto.wait_for_response(), timeout=self._timeout ) return {"id": request_id, "ret": "ok", "resp": payload} except asyncio.TimeoutError: @@ -110,9 +127,9 @@ class MQTTHelperBot: async def send_command( self, cmdjson: dict[str, Any], request_id: str ) -> dict[str, Any]: - if self.client is None: + if self._client is None: await self.start() - assert self.client is not None + assert self._client is not None try: topic = "iot/p2p/{}/helperbot/bumper/helperbot/{}/{}/{}/q/{}/{}".format( @@ -124,14 +141,14 @@ class MQTTHelperBot: cmdjson["payloadType"], ) command_dto = CommandDto(cmdjson["payloadType"]) - self.commands[request_id] = command_dto + self._commands[request_id] = command_dto if cmdjson["payloadType"] == "j": payload = json.dumps(cmdjson["payload"]) else: payload = str(cmdjson["payload"]) - await self.client.publish(topic, payload.encode(), QOS_0) + await self._client.publish(topic, payload.encode(), QOS_0) resp = await self._wait_for_resp(command_dto, request_id) return resp @@ -144,7 +161,16 @@ class MQTTHelperBot: "debug": "exception occurred please check bumper logs", } finally: - self.commands.pop(request_id, None) + self._commands.pop(request_id, None) + + async def disconnect(self): + if self._new_messages_task is not None: + self._new_messages_task.cancel() + self._new_messages_task = None + + if self._client is not None: + await self._client.disconnect() + self._client = None class MQTTServer: @@ -198,6 +224,11 @@ class MQTTServer: except Exception as e: mqttserverlog.exception(f"{e}") + @property + def state(self) -> Broker.states: + """Return the state of the broker.""" + return self._broker.transitions.state + @property def broker(self) -> Broker: return self._broker @@ -206,11 +237,14 @@ class MQTTServer: mqttserverlog.info(f"Starting MQTT Server at {self._host}:{self._port}") try: - await self.broker.start() + await self._broker.start() except Exception as e: mqttserverlog.exception(f"{e}") raise e + async def shutdown(self): + await self._broker.shutdown() + class BumperMQTTServer_Plugin: def __init__(self, context: BrokerContext) -> None: @@ -367,10 +401,6 @@ class BumperMQTTServer_Plugin: helperbotlog.debug( f"Received Response - Topic: {topic} - Message: {data_decoded}" ) - if topic_split[10] in bumper.mqtt_helperbot.commands: - bumper.mqtt_helperbot.commands[topic_split[10]].add_response( - data_decoded - ) elif topic_split[3] == "helperbot": # Helperbot sending command helperbotlog.debug( diff --git a/bumper/web/templates/home.jinja2 b/bumper/web/templates/home.jinja2 index 5af2239..23b1381 100644 --- a/bumper/web/templates/home.jinja2 +++ b/bumper/web/templates/home.jinja2 @@ -148,9 +148,9 @@
- Status: {% if helperbot[0].state == "connected" %} - {{ helperbot[0].state }} {% else %} - {{ helperbot[0].state }} {% endif %} + Status: {% if helperbot.connected %} + Connected {% else %} + Disconnected {% endif %}
diff --git a/tests/conftest.py b/tests/conftest.py index bcf9903..dc6ee31 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,4 +1,7 @@ +import asyncio + import pytest +from amqtt.client import MQTTClient import bumper from tests import HOST, MQTT_PORT @@ -9,10 +12,27 @@ async def mqtt_server(): mqtt_server = bumper.MQTTServer(HOST, MQTT_PORT, password_file="tests/passwd") await mqtt_server.start() bumper.mqtt_server = mqtt_server + while not mqtt_server.state == "started": + await asyncio.sleep(0.1) yield - await mqtt_server.broker.shutdown() + await mqtt_server.shutdown() + + +@pytest.mark.usefixtures("mqtt_server") +@pytest.fixture +async def mqtt_client(): + client = MQTTClient( + client_id="helperbot@bumper/test", + config={"check_hostname": False, "auto_reconnect": False}, + ) + + await client.connect(f"mqtts://{HOST}:{MQTT_PORT}/", cafile=bumper.ca_cert) + + yield client + + await client.disconnect() @pytest.fixture diff --git a/tests/test_confserver.py b/tests/test_confserver.py index af0ee4b..dfcaaff 100644 --- a/tests/test_confserver.py +++ b/tests/test_confserver.py @@ -70,7 +70,7 @@ async def test_base(conf_server_client): resp = await conf_server_client.get("/") assert resp.status == 200 - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() bumper.xmpp_server.disconnect() @@ -100,7 +100,7 @@ async def test_restartService(conf_server_client): resp = await conf_server_client.get("/restart_XMPPServer") assert resp.status == 200 - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() xmpp_server.disconnect() diff --git a/tests/test_init.py b/tests/test_init.py index b9a57b7..a1badb2 100644 --- a/tests/test_init.py +++ b/tests/test_init.py @@ -53,7 +53,7 @@ async def test_start_stop_debug(): asyncio.create_task(b.start()) await asyncio.sleep(0.1) - while b.mqtt_server.broker.transitions.state == "starting": + while b.mqtt_server.state == "starting": await asyncio.sleep(0.1) l.check_present(("bumper", "INFO", "Starting Bumper")) l.clear() diff --git a/tests/test_mqttserver.py b/tests/test_mqttserver.py index 22700e2..662b12a 100644 --- a/tests/test_mqttserver.py +++ b/tests/test_mqttserver.py @@ -13,18 +13,16 @@ from tests import HOST, MQTT_PORT @pytest.mark.usefixtures("mqtt_server") -async def test_helperbot_message(): +async def test_helperbot_message(mqtt_client: MQTTClient): with LogCapture() as l: # Test broadcast message mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT) await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected + assert mqtt_helperbot.is_connected msg_payload = "" msg_topic_name = "iot/atr/DustCaseST/bot_serial/ls1ok3/wC3g/x" - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) await asyncio.sleep(0.1) @@ -36,17 +34,15 @@ async def test_helperbot_message(): ) ) # Check broadcast message was logged l.clear() - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() # Send command to bot mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT) await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected + assert mqtt_helperbot.is_connected msg_payload = "{}" msg_topic_name = "iot/p2p/GetWKVer/helperbot/bumper/helperbot/bot_serial/ls1ok3/wC3g/q/iCmuqp/j" - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) await asyncio.sleep(0.1) @@ -58,17 +54,15 @@ async def test_helperbot_message(): ) ) # Check send command message was logged l.clear() - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() # Received response to command mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT) await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected + assert mqtt_helperbot.is_connected msg_payload = '{"ret":"ok","ver":"0.13.5"}' msg_topic_name = "iot/p2p/GetWKVer/bot_serial/ls1ok3/wC3g/helperbot/bumper/helperbot/p/iCmuqp/j" - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) await asyncio.sleep(0.1) @@ -80,17 +74,15 @@ async def test_helperbot_message(): ) ) # Check received response message was logged l.clear() - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() # Received unknown message mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT) await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected + assert mqtt_helperbot.is_connected msg_payload = "test" msg_topic_name = "iot/p2p/GetWKVer/bot_serial/ls1ok3/wC3g/TESTBAD/bumper/helperbot/p/iCmuqp/j" - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) await asyncio.sleep(0.1) @@ -102,17 +94,15 @@ async def test_helperbot_message(): ) ) # Check received message was logged l.clear() - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() # Received error message mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT) await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected + assert mqtt_helperbot.is_connected msg_payload = "" msg_topic_name = "iot/atr/errors/bot_serial/ls1ok3/wC3g/x" - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) await asyncio.sleep(0.1) @@ -124,19 +114,17 @@ async def test_helperbot_message(): ) ) # Check received message was logged l.clear() - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() @pytest.mark.usefixtures("mqtt_server") -async def test_helperbot_expire_message(): +async def test_helperbot_expire_message(mqtt_client: MQTTClient): timeout = 0.1 # Test broadcast message mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT, timeout) bumper.mqtt_helperbot = mqtt_helperbot await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected + assert mqtt_helperbot.is_connected expire_msg_payload = '{"ret":"ok","ver":"0.13.5"}' expire_msg_topic_name = "iot/p2p/GetWKVer/bot_serial/ls1ok3/wC3g/helperbot/bumper/helperbot/p/testgood/j" @@ -148,33 +136,31 @@ async def test_helperbot_expire_message(): "payload": expire_msg_payload, } - mqtt_helperbot.commands[request_id] = data + mqtt_helperbot._commands[request_id] = data - assert mqtt_helperbot.commands[request_id] == data + assert mqtt_helperbot._commands[request_id] == data await asyncio.sleep(0.1) msg_payload = "" msg_topic_name = "iot/atr/DustCaseST/bot_serial/ls1ok3/wC3g/x" - await mqtt_helperbot.client.publish( + await mqtt_client.publish( msg_topic_name, msg_payload.encode(), QOS_0 ) # Send another message to force get_msg await asyncio.sleep(timeout * 2) - assert mqtt_helperbot.commands.get(request_id, None) == None + assert mqtt_helperbot._commands.get(request_id, None) == None - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() @pytest.mark.usefixtures("mqtt_server") -async def test_helperbot_sendcommand(): +async def test_helperbot_sendcommand(mqtt_client: MQTTClient): timeout = 0.1 mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT, timeout) bumper.mqtt_helperbot = mqtt_helperbot await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected + assert mqtt_helperbot.is_connected cmdjson = { "toType": "ls1ok3", @@ -204,7 +190,7 @@ async def test_helperbot_sendcommand(): # Send response beforehand msg_payload = '{"ret":"ok","ver":"0.13.5"}' msg_topic_name = "iot/p2p/GetWKVer/bot_serial/ls1ok3/wC3g/helperbot/bumper/helperbot/p/testgood/j" - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) commandresult = await mqtt_helperbot.send_command(cmdjson, "testgood") assert commandresult == { @@ -236,7 +222,7 @@ async def test_helperbot_sendcommand(): # Send response beforehand msg_payload = "" msg_topic_name = "iot/p2p/GetLifeSpan/bot_serial/ls1ok3/wC3g/helperbot/bumper/helperbot/p/testx/q" - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) commandresult = await mqtt_helperbot.send_command(cmdjson, "testx") assert commandresult == { @@ -271,7 +257,7 @@ async def test_helperbot_sendcommand(): msg_topic_name = ( "iot/p2p/getStats/bot_serial/ls1ok3/wC3g/helperbot/bumper/helperbot/p/testj/j" ) - await mqtt_helperbot.client.publish(msg_topic_name, msg_payload.encode(), QOS_0) + await mqtt_client.publish(msg_topic_name, msg_payload.encode(), QOS_0) commandresult = await mqtt_helperbot.send_command(cmdjson, "testj") @@ -301,7 +287,7 @@ async def test_helperbot_sendcommand(): "ret": "ok", } - await mqtt_helperbot.client.disconnect() + await mqtt_helperbot.disconnect() async def test_mqttserver(): @@ -320,10 +306,8 @@ async def test_mqttserver(): # Test helperbot connect mqtt_helperbot = bumper.MQTTHelperBot(HOST, MQTT_PORT) await mqtt_helperbot.start() - assert ( - mqtt_helperbot.client._connected_state._value == True - ) # Check helperbot is connected - await mqtt_helperbot.client.disconnect() + assert mqtt_helperbot.is_connected + await mqtt_helperbot.disconnect() # Test client connect bumper.user_add("user_123") # Add user to db @@ -408,7 +392,7 @@ async def test_mqttserver(): order_matters=False, ) finally: - await mqtt_server.broker.shutdown() + await mqtt_server.shutdown() async def test_nofileauth_mqttserver(): @@ -418,7 +402,7 @@ async def test_nofileauth_mqttserver(): HOST, MQTT_PORT, password_file="tests/passwd-notfound" ) await mqtt_server.start() - await mqtt_server.broker.shutdown() + await mqtt_server.shutdown() l.check_present( (