Skip to content

Commit 36e498b

Browse files
committed
Added additional tests
1 parent 404b6a3 commit 36e498b

10 files changed

Lines changed: 1210 additions & 69 deletions

File tree

tb_mqtt_client/common/config_loader.py

Lines changed: 21 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,17 @@ class DeviceConfig:
2424
"""
2525

2626
def __init__(self, config=None):
27+
self.host = None
28+
self.port = 1883
29+
self.access_token: Optional[str] = None
30+
self.username: Optional[str] = None
31+
self.password: Optional[str] = None
32+
self.client_id: Optional[str] = None
33+
self.ca_cert: Optional[str] = None
34+
self.client_cert: Optional[str] = None
35+
self.private_key: Optional[str] = None
36+
self.qos: int = 1
37+
2738
if config is not None:
2839
self.host: str = config.get("host", "localhost")
2940
self.port: int = config.get("port", 1883)
@@ -35,24 +46,24 @@ def __init__(self, config=None):
3546
self.client_cert: Optional[str] = config.get("client_cert")
3647
self.private_key: Optional[str] = config.get("private_key")
3748

38-
self.host: str = os.getenv("TB_HOST")
39-
self.port: int = int(os.getenv("TB_PORT", 1883))
49+
self.host: str = os.getenv("TB_HOST", self.host)
50+
self.port: int = int(os.getenv("TB_PORT", self.port))
4051

4152
# Authentication options
42-
self.access_token: Optional[str] = os.getenv("TB_ACCESS_TOKEN")
43-
self.username: Optional[str] = os.getenv("TB_USERNAME")
44-
self.password: Optional[str] = os.getenv("TB_PASSWORD")
53+
self.access_token: Optional[str] = os.getenv("TB_ACCESS_TOKEN", self.access_token)
54+
self.username: Optional[str] = os.getenv("TB_USERNAME", self.username)
55+
self.password: Optional[str] = os.getenv("TB_PASSWORD", self.password)
4556

4657
# Optional
47-
self.client_id: Optional[str] = os.getenv("TB_CLIENT_ID")
58+
self.client_id: Optional[str] = os.getenv("TB_CLIENT_ID", self.client_id)
4859

4960
# TLS options
50-
self.ca_cert: Optional[str] = os.getenv("TB_CA_CERT")
51-
self.client_cert: Optional[str] = os.getenv("TB_CLIENT_CERT")
52-
self.private_key: Optional[str] = os.getenv("TB_PRIVATE_KEY")
61+
self.ca_cert: Optional[str] = os.getenv("TB_CA_CERT", self.ca_cert)
62+
self.client_cert: Optional[str] = os.getenv("TB_CLIENT_CERT", self.client_cert)
63+
self.private_key: Optional[str] = os.getenv("TB_PRIVATE_KEY", self.private_key)
5364

5465
# Default values
55-
self.qos: int = int(os.getenv("TB_QOS", 1))
66+
self.qos: int = int(os.getenv("TB_QOS", self.qos))
5667

5768
def use_tls_auth(self) -> bool:
5869
return all([self.ca_cert, self.client_cert, self.private_key])

tb_mqtt_client/common/gmqtt_patch.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -122,8 +122,8 @@ def parse_mqtt_properties(packet: bytes) -> dict:
122122
properties_dict = defaultdict(list)
123123

124124
try:
125-
properties_len, _ = unpack_variable_byte_integer(packet)
126-
props = packet[:properties_len]
125+
properties_len, rest = unpack_variable_byte_integer(packet)
126+
props = rest[:properties_len] # slice out exactly the properties section
127127

128128
while props:
129129
property_identifier = props[0]

tb_mqtt_client/common/rate_limit/backpressure_controller.py

Lines changed: 0 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,6 @@ def __init__(self, main_stop_event: Event):
3030
self._consecutive_quota_exceeded = 0
3131
self._last_quota_exceeded = datetime.now(UTC)
3232
self._max_backoff_seconds = 3600 # 1 hour
33-
self._can_process_messages_events: List[asyncio.Event] = []
3433
logger.debug("BackpressureController initialized with default pause duration of %s seconds",
3534
self._default_pause_duration.total_seconds())
3635

@@ -90,25 +89,10 @@ def should_pause(self) -> bool:
9089
# Reset pause state
9190
self._pause_until = None
9291
logger.info("Backpressure released, resuming publishing")
93-
for event in self._can_process_messages_events:
94-
if not event.is_set():
95-
event.set()
96-
logger.debug("Set can-process event %s", event)
9792
return False
9893

9994
def clear(self):
10095
if self._pause_until is not None:
10196
logger.info("Clearing backpressure pause")
10297
self._pause_until = None
10398
self._consecutive_quota_exceeded = 0
104-
105-
def register_can_process_event(self, event: Event):
106-
"""
107-
Register an event that will be set when the controller can process messages again.
108-
This is useful for other components to wait until the backpressure is lifted.
109-
"""
110-
if not isinstance(event, Event):
111-
raise ValueError("Expected an asyncio.Event instance")
112-
self._can_process_messages_events.append(event)
113-
logger.debug("Registered a new can-process event, total events: %d",
114-
len(self._can_process_messages_events))

tests/common/test_async_utils.py

Lines changed: 55 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,13 +13,36 @@
1313
# limitations under the License.
1414

1515
import asyncio
16+
import threading
17+
import time
1618

1719
import pytest
1820

19-
from tb_mqtt_client.common.async_utils import FutureMap, await_or_stop
21+
import tb_mqtt_client.common.async_utils as async_utils_mod
22+
from tb_mqtt_client.common.async_utils import FutureMap, await_or_stop, run_coroutine_sync
2023
from tb_mqtt_client.common.publish_result import PublishResult
2124

2225

26+
@pytest.fixture
27+
def fake_loop(monkeypatch):
28+
class _FakeTask:
29+
def __init__(self, thread: threading.Thread):
30+
self._thread = thread
31+
32+
def done(self): # optional helpers if ever needed
33+
return not self._thread.is_alive()
34+
35+
class _FakeLoop:
36+
def create_task(self, coro):
37+
t = threading.Thread(target=lambda: asyncio.run(coro), daemon=True)
38+
t.start()
39+
return _FakeTask(t)
40+
41+
monkeypatch.setattr(async_utils_mod.asyncio, "get_running_loop", lambda: _FakeLoop())
42+
43+
return _FakeLoop()
44+
45+
2346
@pytest.mark.asyncio
2447
async def test_future_map_register_and_get_parents():
2548
fm = FutureMap()
@@ -150,5 +173,36 @@ async def coro():
150173
assert result is None
151174

152175

176+
def test_returns_result_when_coroutine_completes(fake_loop):
177+
async def ok_coro():
178+
await asyncio.sleep(0.01)
179+
return "done"
180+
181+
result = run_coroutine_sync(lambda: ok_coro(), timeout=0.5)
182+
assert result == "done"
183+
184+
185+
def test_raises_original_exception_from_coroutine(fake_loop):
186+
class CustomError(RuntimeError):
187+
pass
188+
189+
async def bad_coro():
190+
await asyncio.sleep(0.01)
191+
raise CustomError("boom")
192+
193+
with pytest.raises(CustomError, match="boom"):
194+
run_coroutine_sync(lambda: bad_coro(), timeout=0.5)
195+
196+
197+
def test_timeout_raises_timeout_error(fake_loop):
198+
async def slow_coro():
199+
await asyncio.sleep(0.2)
200+
return "too late"
201+
202+
with pytest.raises(TimeoutError, match=r"did not complete in 0\.05 seconds"):
203+
run_coroutine_sync(lambda: slow_coro(), timeout=0.05, raise_on_timeout=True)
204+
time.sleep(0.25)
205+
206+
153207
if __name__ == '__main__':
154208
pytest.main([__file__, "--tb=short", "-v"])

tests/common/test_config_loader.py

Lines changed: 35 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,32 @@
2020

2121
class TestDeviceConfig(unittest.TestCase):
2222

23-
def loads_default_values_when_env_vars_missing(self):
23+
def test_config_creation_from_dict(self):
24+
config_dict = {
25+
"host": "localhost",
26+
"port": 1883,
27+
"access_token": "test_token",
28+
"username": "test_user",
29+
"password": "test_pass",
30+
"client_id": "test_client",
31+
"ca_cert": "test_ca",
32+
"client_cert": "test_cert",
33+
"private_key": "test_key",
34+
"qos": 1
35+
}
36+
config = DeviceConfig(config_dict)
37+
self.assertEqual(config.host, "localhost")
38+
self.assertEqual(config.port, 1883)
39+
self.assertEqual(config.access_token, "test_token")
40+
self.assertEqual(config.username, "test_user")
41+
self.assertEqual(config.password, "test_pass")
42+
self.assertEqual(config.client_id, "test_client")
43+
self.assertEqual(config.ca_cert, "test_ca")
44+
self.assertEqual(config.client_cert, "test_cert")
45+
self.assertEqual(config.private_key, "test_key")
46+
self.assertEqual(config.qos, 1)
47+
48+
def test_loads_default_values_when_env_vars_missing(self):
2449
os.environ.clear()
2550
config = DeviceConfig()
2651
self.assertEqual(config.host, None)
@@ -34,7 +59,7 @@ def loads_default_values_when_env_vars_missing(self):
3459
self.assertEqual(config.private_key, None)
3560
self.assertEqual(config.qos, 1)
3661

37-
def loads_values_from_env_vars(self):
62+
def test_loads_values_from_env_vars(self):
3863
os.environ["TB_HOST"] = "test_host"
3964
os.environ["TB_PORT"] = "8883"
4065
os.environ["TB_ACCESS_TOKEN"] = "test_token"
@@ -58,22 +83,22 @@ def loads_values_from_env_vars(self):
5883
self.assertEqual(config.private_key, "test_key")
5984
self.assertEqual(config.qos, 2)
6085

61-
def detects_tls_auth_correctly(self):
86+
def test_detects_tls_auth_correctly(self):
6287
os.environ["TB_CA_CERT"] = "test_ca"
6388
os.environ["TB_CLIENT_CERT"] = "test_cert"
6489
os.environ["TB_PRIVATE_KEY"] = "test_key"
6590
config = DeviceConfig()
6691
self.assertTrue(config.use_tls_auth())
6792

68-
def detects_tls_correctly(self):
93+
def test_detects_tls_correctly(self):
6994
os.environ["TB_CA_CERT"] = "test_ca"
7095
config = DeviceConfig()
7196
self.assertTrue(config.use_tls())
7297

7398

7499
class TestGatewayConfig(unittest.TestCase):
75100

76-
def loads_gateway_specific_env_vars(self):
101+
def test_loads_gateway_specific_env_vars(self):
77102
os.environ["TB_GW_HOST"] = "gw_host"
78103
os.environ["TB_GW_PORT"] = "8884"
79104
os.environ["TB_GW_ACCESS_TOKEN"] = "gw_token"
@@ -97,10 +122,14 @@ def loads_gateway_specific_env_vars(self):
97122
self.assertEqual(config.private_key, "gw_key")
98123
self.assertEqual(config.qos, 0)
99124

100-
def falls_back_to_device_config_when_gateway_env_vars_missing(self):
125+
def test_falls_back_to_device_config_when_gateway_env_vars_missing(self):
101126
os.environ.clear()
102127
os.environ["TB_HOST"] = "device_host"
103128
os.environ["TB_PORT"] = "1884"
104129
config = GatewayConfig()
105130
self.assertEqual(config.host, "device_host")
106131
self.assertEqual(config.port, 1884)
132+
133+
134+
if __name__ == "__main__":
135+
unittest.main()

0 commit comments

Comments
 (0)