import pytest import sys import os from unittest.mock import MagicMock, patch sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from modules.mqtt_client import MQTTClient def _setup_mock_client(mock_client_class): """创建 mock 客户端,并让 loop_start 触发 on_connect 回调""" mock_client = MagicMock() mock_client.connect.return_value = 0 mock_client_class.return_value = mock_client # 捕获 on_connect 回调,在 loop_start 时触发它 def fake_loop_start(): import threading, time # 延迟调用以模拟异步回调,等待 connect() 完成 self.config 赋值 def _trigger(): time.sleep(0.05) mock_client.on_connect(mock_client, None, {}, 0) t = threading.Thread(target=_trigger, daemon=True) t.start() mock_client.loop_start.side_effect = fake_loop_start return mock_client class TestMQTTClient: def setup_method(self): self.mqtt_client = MQTTClient() self.data_callback = MagicMock() self.status_callback = MagicMock() self.mqtt_client.set_data_callback(self.data_callback) self.mqtt_client.set_status_callback(self.status_callback) def teardown_method(self): self.mqtt_client.disconnect() @patch('modules.mqtt_client.mqtt.Client') def test_connect_success(self, mock_client_class): """测试成功连接MQTT服务器""" mock_client = _setup_mock_client(mock_client_class) success, message = self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) assert success is True assert 'localhost' in message mock_client.connect.assert_called_once() mock_client.loop_start.assert_called_once() @patch('modules.mqtt_client.mqtt.Client') def test_connect_with_tls(self, mock_client_class): """测试使用TLS连接MQTT服务器""" mock_client = _setup_mock_client(mock_client_class) success, message = self.mqtt_client.connect( broker='localhost', port=8883, client_id='test_client', keepalive=60, use_tls=True, ca_certs='path/to/ca.crt' ) assert success is True mock_client.tls_set.assert_called_once() @patch('modules.mqtt_client.mqtt.Client') def test_connect_with_auth(self, mock_client_class): """测试使用用户名密码连接MQTT服务器""" mock_client = _setup_mock_client(mock_client_class) success, message = self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60, username='test_user', password='test_pass' ) assert success is True mock_client.username_pw_set.assert_called_once_with('test_user', 'test_pass') @patch('modules.mqtt_client.mqtt.Client') def test_disconnect(self, mock_client_class): """测试断开MQTT连接""" mock_client = _setup_mock_client(mock_client_class) self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) self.mqtt_client.disconnect() mock_client.loop_stop.assert_called_once() mock_client.disconnect.assert_called_once() @patch('modules.mqtt_client.mqtt.Client') def test_publish_success(self, mock_client_class): """测试成功发布消息""" import paho.mqtt.client as paho mock_client = _setup_mock_client(mock_client_class) mock_publish = MagicMock() mock_publish.rc = paho.MQTT_ERR_SUCCESS mock_client.publish.return_value = mock_publish self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) success, message = self.mqtt_client.publish('test/topic', 'test message') assert success is True mock_client.publish.assert_called_once() @patch('modules.mqtt_client.mqtt.Client') def test_publish_not_connected(self, mock_client_class): """测试未连接时发布消息""" success, message = self.mqtt_client.publish('test/topic', 'test message') assert success is False assert '未连接' in message @patch('modules.mqtt_client.mqtt.Client') def test_subscribe_success(self, mock_client_class): """测试成功订阅主题""" import paho.mqtt.client as paho mock_client = _setup_mock_client(mock_client_class) mock_client.subscribe.return_value = (paho.MQTT_ERR_SUCCESS, 1) self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) success, message = self.mqtt_client.subscribe('test/topic') assert success is True mock_client.subscribe.assert_called_once() @patch('modules.mqtt_client.mqtt.Client') def test_subscribe_failure(self, mock_client_class): """测试订阅主题失败""" mock_client = _setup_mock_client(mock_client_class) mock_client.subscribe.return_value = (1, None) self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) success, message = self.mqtt_client.subscribe('test/topic') assert success is False @patch('modules.mqtt_client.mqtt.Client') def test_unsubscribe(self, mock_client_class): """测试取消订阅主题""" import paho.mqtt.client as paho mock_client = _setup_mock_client(mock_client_class) mock_unsub = MagicMock() mock_unsub.rc = paho.MQTT_ERR_SUCCESS mock_client.unsubscribe.return_value = mock_unsub self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) # subscribe first so topics list is populated mock_sub = MagicMock() mock_sub.rc = paho.MQTT_ERR_SUCCESS mock_client.subscribe.return_value = mock_sub self.mqtt_client.subscribe('test/topic') success, message = self.mqtt_client.unsubscribe('test/topic') assert success is True mock_client.unsubscribe.assert_called_once() @patch('modules.mqtt_client.mqtt.Client') def test_get_status(self, mock_client_class): """测试获取MQTT客户端状态""" mock_client = _setup_mock_client(mock_client_class) # 初始状态 st = self.mqtt_client.get_status() assert isinstance(st, dict) assert st['connected'] is False # 连接后 self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) st = self.mqtt_client.get_status() assert isinstance(st, dict) assert st['connected'] is True # 断开后 self.mqtt_client.disconnect() st = self.mqtt_client.get_status() assert st['connected'] is False @patch('modules.mqtt_client.mqtt.Client') def test_on_message_callback(self, mock_client_class): """测试MQTT消息回调处理""" mock_client = _setup_mock_client(mock_client_class) self.mqtt_client.connect( broker='localhost', port=1883, client_id='test_client', keepalive=60 ) on_message_callback = mock_client.on_message mock_message = MagicMock() mock_message.topic = "test/topic" mock_message.payload = b"test payload" mock_message.qos = 0 mock_message.retain = False on_message_callback(mock_client, None, mock_message) self.data_callback.assert_called_once() call_data = self.data_callback.call_args[0][0] assert call_data['topic'] == 'test/topic' assert call_data['payload'] == 'test payload' if __name__ == "__main__": pytest.main(["-v", __file__])