Skip to content

Commit d841be2

Browse files
committed
fix: cancel non-resumable MQTT operations
1 parent 07707fe commit d841be2

4 files changed

Lines changed: 164 additions & 65 deletions

File tree

‎azure-iot-device/azure/iot/device/common/mqtt_transport.py‎

Lines changed: 34 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -726,7 +726,9 @@ def disconnect(self, clear_inflight=False):
726726
"""
727727
Disconnect from the MQTT Server and wait for the network loop to stop.
728728
729-
Optionally, clear any inflight operation tracking if clear_inflight is True.
729+
If clear_inflight is True, complete all tracked operations as cancelled. Otherwise,
730+
preserve resumable QoS 1 and QoS 2 publishes and complete non-resumable operations as
731+
cancelled.
730732
731733
:raises: ProtocolClientError if there is some client error.
732734
:raises: ConnectionDroppedError in unexpected cases.
@@ -759,7 +761,7 @@ def disconnect(self, clear_inflight=False):
759761
if clear_inflight:
760762
self._op_manager.complete_all_tracked_operations_as_cancelled()
761763
else:
762-
self._op_manager.stop_tracking_non_publish_operations()
764+
self._op_manager.complete_non_resumable_operations_as_cancelled()
763765
if self._connection_lifecycle:
764766
self._connection_lifecycle.finish_disconnect()
765767
else:
@@ -774,7 +776,7 @@ def disconnect(self, clear_inflight=False):
774776
if clear_inflight:
775777
self._op_manager.complete_all_tracked_operations_as_cancelled()
776778
else:
777-
self._op_manager.stop_tracking_non_publish_operations()
779+
self._op_manager.complete_non_resumable_operations_as_cancelled()
778780
if self._connection_lifecycle:
779781
self._connection_lifecycle.finish_disconnect()
780782

@@ -893,12 +895,20 @@ def publish(self, topic, payload, qos=1, callback=None):
893895
"Paho retained QoS {} PUBLISH with MID {} for the next connection".format(qos, mid)
894896
)
895897
self._op_manager.register_operation(
896-
mid=mid, callback=callback, operation_type=OperationType.PUBLISH
898+
mid=mid,
899+
callback=callback,
900+
operation_type={
901+
0: OperationType.PUBLISH_QOS_0,
902+
1: OperationType.PUBLISH_QOS_1,
903+
2: OperationType.PUBLISH_QOS_2,
904+
}[qos],
897905
)
898906

899907

900908
class OperationType(Enum):
901-
PUBLISH = "PUBLISH"
909+
PUBLISH_QOS_0 = "PUBLISH_QOS_0"
910+
PUBLISH_QOS_1 = "PUBLISH_QOS_1"
911+
PUBLISH_QOS_2 = "PUBLISH_QOS_2"
902912
SUBSCRIBE = "SUBSCRIBE"
903913
UNSUBSCRIBE = "UNSUBSCRIBE"
904914

@@ -1073,22 +1083,26 @@ def complete_operation(self, mid, error=None):
10731083
# Completion callbacks are optional.
10741084
logger.debug("No callback set for Paho MID {}".format(mid))
10751085

1076-
def stop_tracking_non_publish_operations(self):
1077-
"""Stop tracking SUBSCRIBE and UNSUBSCRIBE operations without invoking callbacks.
1086+
def complete_non_resumable_operations_as_cancelled(self):
1087+
"""Complete operations Paho cannot resume as cancelled.
10781088
1079-
Paho does not retain these operations for a later connection. PUBLISH operations remain
1080-
tracked because Paho owns their MQTT 3.1.1 QoS retransmission state.
1089+
Paho retains MQTT 3.1.1 QoS 1 and QoS 2 PUBLISH operations for a later connection.
1090+
SUBSCRIBE, UNSUBSCRIBE, and QoS 0 PUBLISH operations are not retained by Paho, so remove
1091+
their tracking, tombstone their MIDs, and invoke their callbacks with ``cancelled=True``.
10811092
"""
1093+
logger.debug("Completing non-resumable tracked operations as cancelled")
10821094
with self._lock:
1083-
matching_mids = [
1084-
mid
1095+
pending_ops = [
1096+
(mid, pending_operation)
10851097
for mid, pending_operation in self._pending_operations.items()
10861098
if pending_operation.operation_type
1087-
in (OperationType.SUBSCRIBE, OperationType.UNSUBSCRIBE)
1099+
not in (OperationType.PUBLISH_QOS_1, OperationType.PUBLISH_QOS_2)
10881100
]
1089-
for mid in matching_mids:
1101+
for mid, _ in pending_ops:
10901102
del self._pending_operations[mid]
1091-
self._cancelled_operation_mids.update(matching_mids)
1103+
self._cancelled_operation_mids.update(mid for mid, _ in pending_ops)
1104+
1105+
self._defer_or_invoke_cancellation_callbacks(pending_ops)
10921106

10931107
def complete_all_tracked_operations_as_cancelled(self):
10941108
"""Complete all tracked SDK operations as cancelled and clear unknown completions.
@@ -1107,13 +1121,16 @@ def complete_all_tracked_operations_as_cancelled(self):
11071121
self._pending_operations.clear()
11081122
self._unknown_operation_completions.clear()
11091123

1124+
self._defer_or_invoke_cancellation_callbacks(pending_ops)
1125+
1126+
def _defer_or_invoke_cancellation_callbacks(self, pending_operations):
11101127
deferred_operations = getattr(
11111128
self._deferred_cancellation_callbacks, "pending_operations", None
11121129
)
11131130
if deferred_operations is not None:
1114-
deferred_operations.extend(pending_ops)
1131+
deferred_operations.extend(pending_operations)
11151132
else:
1116-
self._invoke_cancellation_callbacks(pending_ops)
1133+
self._invoke_cancellation_callbacks(pending_operations)
11171134

11181135
def _invoke_cancellation_callbacks(self, pending_operations):
11191136
"""Invoke callbacks for operations whose tracking was cancelled."""
@@ -1140,6 +1157,7 @@ def _invoke_cancellation_callbacks(self, pending_operations):
11401157

11411158
# TODO: Clarify hard-disconnect semantics because cancelling an SDK publish operation does not
11421159
# prevent Paho from delivering a retained QoS 1 or QoS 2 message after a later connection.
1160+
# Re-evaluate the inclusion of "hard" disconnect.
11431161

11441162
# NOTE: Connection lifecycle calls are deliberately serialized here and by ConnectionStateStage.
11451163
# CONNECTION_TIMEOUT bounds the wait for CONNACK, allowing queued lifecycle operations such as

‎azure-iot-device/azure/iot/device/common/pipeline/pipeline_stages_mqtt.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -330,11 +330,11 @@ def _reconcile_mqtt_operation_tracking_after_connection_drop(self):
330330
self.transport._op_manager.complete_all_tracked_operations_as_cancelled()
331331
else:
332332
logger.debug(
333-
"{}: Connection Retry enabled - preserving PUBLISH tracking and stopping SUBSCRIBE and UNSUBSCRIBE tracking".format(
333+
"{}: Connection Retry enabled - preserving resumable PUBLISH tracking and completing non-resumable MQTT operations as cancelled".format(
334334
self.name
335335
)
336336
)
337-
self.transport._op_manager.stop_tracking_non_publish_operations()
337+
self.transport._op_manager.complete_non_resumable_operations_as_cancelled()
338338

339339
@pipeline_thread.runs_on_pipeline_thread
340340
def _complete_pending_connection_op_after_disconnect(self, cause=None):

‎tests/unit/common/pipeline/test_pipeline_stages_mqtt.py‎

Lines changed: 20 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -645,7 +645,10 @@ def test_does_not_apply_connection_drop_handling(self, stage, op, fake_pipeline_
645645
assert (
646646
stage.transport._op_manager.complete_all_tracked_operations_as_cancelled.call_count == 0
647647
)
648-
assert stage.transport._op_manager.stop_tracking_non_publish_operations.call_count == 0
648+
assert (
649+
stage.transport._op_manager.complete_non_resumable_operations_as_cancelled.call_count
650+
== 0
651+
)
649652
assert stage.report_background_exception.call_count == 0
650653

651654
@pytest.mark.it(
@@ -966,13 +969,19 @@ def test_connection_drop_handling(
966969
stage.transport._op_manager.complete_all_tracked_operations_as_cancelled.call_count
967970
== 0
968971
)
969-
assert stage.transport._op_manager.stop_tracking_non_publish_operations.call_count == 1
972+
assert (
973+
stage.transport._op_manager.complete_non_resumable_operations_as_cancelled.call_count
974+
== 1
975+
)
970976
else:
971977
assert (
972978
stage.transport._op_manager.complete_all_tracked_operations_as_cancelled.call_count
973979
== 1
974980
)
975-
assert stage.transport._op_manager.stop_tracking_non_publish_operations.call_count == 0
981+
assert (
982+
stage.transport._op_manager.complete_non_resumable_operations_as_cancelled.call_count
983+
== 0
984+
)
976985

977986
assert stage.report_background_exception.call_count == 1
978987
background_exception = stage.report_background_exception.call_args.args[0]
@@ -992,7 +1001,7 @@ def test_connection_drop_handling(
9921001
),
9931002
pytest.param(
9941003
True,
995-
"stop_tracking_non_publish_operations",
1004+
"complete_non_resumable_operations_as_cancelled",
9961005
id="Connection retry enabled",
9971006
),
9981007
],
@@ -1084,22 +1093,24 @@ def test_completes_tracked_operations_without_retry(self, mocker, stage, arbitra
10841093
assert mock_cancel.call_args == mocker.call()
10851094

10861095
@pytest.mark.it(
1087-
"Preserves publishes and stops tracking other MQTT operations if connection retry is enabled"
1096+
"Preserves resumable publishes and cancels other MQTT operations if connection retry is enabled"
10881097
)
1089-
def test_preserves_publish_tracking_with_retry(self, mocker, stage, arbitrary_exception):
1098+
def test_cancels_non_resumable_operations_with_retry(self, mocker, stage, arbitrary_exception):
10901099
stage.transport._op_manager = mocker.MagicMock()
10911100
mock_cancel = stage.transport._op_manager.complete_all_tracked_operations_as_cancelled
1092-
mock_stop_non_publish = stage.transport._op_manager.stop_tracking_non_publish_operations
1101+
mock_cancel_non_resumable = (
1102+
stage.transport._op_manager.complete_non_resumable_operations_as_cancelled
1103+
)
10931104
stage.nucleus.pipeline_configuration.connection_retry = True
10941105
assert stage._pending_connection_op is None
10951106
assert mock_cancel.call_count == 0
1096-
assert mock_stop_non_publish.call_count == 0
1107+
assert mock_cancel_non_resumable.call_count == 0
10971108

10981109
# Trigger disconnect
10991110
stage.transport.on_mqtt_connection_dropped_handler(arbitrary_exception)
11001111

11011112
assert mock_cancel.call_count == 0
1102-
assert mock_stop_non_publish.call_args == mocker.call()
1113+
assert mock_cancel_non_resumable.call_args == mocker.call()
11031114

11041115
@pytest.mark.it("Raises a ConnectionDroppedError as a background exception")
11051116
def test_background_exception_raised(self, stage, arbitrary_exception):

0 commit comments

Comments
 (0)