| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226 |
- 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__])
|