From 6efedd3143262d60393d31c80d122ce041344fd9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Fabr=C3=ADcio=20Bracht?= Date: Wed, 23 Sep 2026 01:13:30 -0300 Subject: [PATCH 1/2] fix MQTT v5 client conformance violations in native and wasm clients --- CHANGELOG.md | 63 + crates/mqtt5-conformance/CONFORMANCE_DIARY.md | 14 + .../src/test_client/inprocess.rs | 9 +- crates/mqtt5-protocol/Cargo.toml | 2 +- crates/mqtt5-protocol/src/packet.rs | 39 +- crates/mqtt5-protocol/src/packet_id.rs | 25 + .../mqtt5-protocol/src/session/topic_alias.rs | 41 +- crates/mqtt5-protocol/src/validation/mod.rs | 118 +- crates/mqtt5-wasm/Cargo.toml | 9 +- crates/mqtt5-wasm/examples/README.md | 4 +- .../examples/qos2-recovery/index.html | 1 + .../examples/session-recovery/index.html | 2 + crates/mqtt5-wasm/examples/websocket/app.js | 2 +- crates/mqtt5-wasm/src/client/callbacks.rs | 106 +- crates/mqtt5-wasm/src/client/connection.rs | 233 +++ crates/mqtt5-wasm/src/client/handlers.rs | 533 ++--- crates/mqtt5-wasm/src/client/keepalive.rs | 39 +- crates/mqtt5-wasm/src/client/mod.rs | 741 +++---- crates/mqtt5-wasm/src/client/outbound.rs | 120 ++ crates/mqtt5-wasm/src/client/packet.rs | 41 + crates/mqtt5-wasm/src/client/qos.rs | 220 ++- crates/mqtt5-wasm/src/client/reader.rs | 139 +- crates/mqtt5-wasm/src/client/reconnect.rs | 178 +- crates/mqtt5-wasm/src/client/state.rs | 219 ++- crates/mqtt5-wasm/src/config.rs | 311 +-- crates/mqtt5-wasm/src/decoder.rs | 104 +- .../mqtt5-wasm/src/transport/message_port.rs | 19 +- crates/mqtt5-wasm/src/transport/websocket.rs | 77 +- crates/mqtt5-wasm/tests/conformance_client.rs | 1606 +++++++++++++++ crates/mqtt5/Cargo.toml | 4 +- crates/mqtt5/examples/deferred_ack.rs | 1 + crates/mqtt5/src/broker/bridge/connection.rs | 11 +- crates/mqtt5/src/client/direct/ack.rs | 161 +- crates/mqtt5/src/client/direct/handlers.rs | 145 +- crates/mqtt5/src/client/direct/keepalive.rs | 2 + crates/mqtt5/src/client/direct/mod.rs | 736 ++++--- crates/mqtt5/src/client/direct/outbound.rs | 250 +++ crates/mqtt5/src/client/direct/reader.rs | 243 ++- crates/mqtt5/src/client/direct/replay.rs | 166 ++ crates/mqtt5/src/client/direct/unified.rs | 79 +- crates/mqtt5/src/client/inner.rs | 56 +- crates/mqtt5/src/client/mod.rs | 91 +- crates/mqtt5/src/client/state.rs | 55 +- crates/mqtt5/src/lib.rs | 2 + crates/mqtt5/src/session.rs | 3 - crates/mqtt5/src/session/flow_control.rs | 292 ++- crates/mqtt5/src/session/retained.rs | 224 --- crates/mqtt5/src/session/state.rs | 384 ++-- crates/mqtt5/src/test_utils.rs | 33 - crates/mqtt5/src/transport/packet_io.rs | 130 +- crates/mqtt5/src/transport/websocket.rs | 156 +- crates/mqtt5/src/types.rs | 27 + crates/mqtt5/tests/conf_client_a.rs | 1670 ++++++++++++++++ crates/mqtt5/tests/conf_client_b.rs | 1509 ++++++++++++++ crates/mqtt5/tests/conf_client_c.rs | 961 +++++++++ crates/mqtt5/tests/conf_client_d.rs | 1735 +++++++++++++++++ crates/mqtt5/tests/deferred_ack_matrix.rs | 14 +- .../mqtt5/tests/integration_complete_flow.rs | 101 +- crates/mqtt5/tests/persistence.rs | 274 ++- crates/mqtt5/tests/retained_messages.rs | 228 --- crates/mqtt5/tests/session_security.rs | 24 +- .../tests/session_state_property_tests.rs | 221 +-- .../mqtt5/tests/session_takeover_backlog.rs | 1 + .../tests/subscription_options_persistence.rs | 46 +- crates/mqttv5-cli/Cargo.toml | 4 +- crates/mqttv5-cli/src/commands/client_args.rs | 137 ++ crates/mqttv5-cli/src/commands/mod.rs | 1 + crates/mqttv5-cli/src/commands/pub_cmd.rs | 352 ++-- crates/mqttv5-cli/src/commands/sub_cmd.rs | 231 +-- 69 files changed, 12239 insertions(+), 3536 deletions(-) create mode 100644 crates/mqtt5-wasm/src/client/connection.rs create mode 100644 crates/mqtt5-wasm/src/client/outbound.rs create mode 100644 crates/mqtt5-wasm/tests/conformance_client.rs create mode 100644 crates/mqtt5/src/client/direct/outbound.rs create mode 100644 crates/mqtt5/src/client/direct/replay.rs delete mode 100644 crates/mqtt5/src/session/retained.rs create mode 100644 crates/mqtt5/tests/conf_client_a.rs create mode 100644 crates/mqtt5/tests/conf_client_b.rs create mode 100644 crates/mqtt5/tests/conf_client_c.rs create mode 100644 crates/mqtt5/tests/conf_client_d.rs delete mode 100644 crates/mqtt5/tests/retained_messages.rs create mode 100644 crates/mqttv5-cli/src/commands/client_args.rs diff --git a/CHANGELOG.md b/CHANGELOG.md index 032b0c87..24e75251 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,69 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). +## [mqtt5 0.41.0] - 2026-09-23 + +A client-side conformance audit drove the real `MqttClient` against a raw-byte fake broker for each of the 149 normative statements in MQTT v5.0 that apply to a client. About 40 MUST statements failed. All are fixed here, and each is pinned by a test named after its OASIS statement ID in `crates/mqtt5/tests/conf_client_{a,b,c,d}.rs`. + +### Breaking + +- **A fresh client that receives Session Present=1 now closes the connection** with DISCONNECT 0x82 and `connect` returns an error, as `[MQTT-3.2.2-4]` requires. A client instance counts as holding session state once it has connected before. To resume a broker-held session from a freshly started process on purpose (the deferred-ack crash-recovery pattern), set the new `ConnectOptions::resume_existing_session` / `with_resume_existing_session(true)`. Session Present=1 in answer to Clean Start=1 is always rejected. +- **`ConnectOptions` has a new public field, `resume_existing_session`.** Code that builds `ConnectOptions` with a struct literal must add it. +- **Invalid outbound requests are now rejected before anything is sent.** This covers topic names with wildcards, empty topics without a Topic Alias, malformed topic filters and `$share` filters, wildcard Response Topics, Subscription Identifiers on PUBLISH or 0 on SUBSCRIBE, and Topic Alias 0 or above the server maximum. It also covers RETAIN when the server reports Retain Available=0, wildcard, shared or subscription-identifier subscriptions the server reports unavailable, and SUBSCRIBE/UNSUBSCRIBE over the server Maximum Packet Size. Previously all of these were sent and left for the broker to reject. +- **Removed the deprecated session-level retained message store.** It was scheduled for removal in 0.32.0 and nothing in the client or broker used it. The broker's retained messages (`broker::storage::RetainedMessage`, `MessageRouter::get_retained_messages`) are unaffected. Removed: + - the `mqtt5::session::retained` module + - `RetainedMessageStore` and `RetainedMessage`, and their `mqtt5::session` re-exports + - `SessionState::store_retained_message`, `get_retained_messages` and `retained_messages` + - `test_utils::test_retained_message` + - `test_utils::TestMessageBuilder::build_retained_batch` +- **Removed the deprecated `WebSocketConfig::with_tls_verification` and `WebSocketConfig::verify_tls`.** Nothing read `verify_tls`. Use `tls_config`. + +### Fixed + +- **Unacknowledged QoS 1/2 PUBLISH and PUBREL are resent when a session resumes** (`[MQTT-4.4.0-1]`), with DUP=1, their original packet identifiers and in their original order (`[MQTT-4.6.0-1]`, `[MQTT-4.6.0-4]`), and within the new connection's Receive Maximum. Before, they were stored and never resent, so QoS 1/2 delivery did not survive a reconnect. QoS 1 entries are now also removed on PUBACK. They used to accumulate forever. +- **Packet identifiers are no longer reused while in flight** (`[MQTT-2.2.1-3]`, `[MQTT-4.3.2-1]`). Allocation skips identifiers held by unacknowledged PUBLISH, outstanding PUBREL and pending SUBSCRIBE/UNSUBSCRIBE. +- **The offline queue goes through the normal publish path.** It used to bypass the send quota, unacknowledged-message tracking, Maximum QoS and Retain Available, and it set DUP=1 on a message's first transmission (`[MQTT-4.3.2-2]`, `[MQTT-3.3.4-7]`, `[MQTT-3.2.2-11]`, `[MQTT-3.2.2-14]`). +- **Send quota is reset on every connection** (`[MQTT-4.9.0-1]`). A publish that timed out waiting for its ack no longer leaks its Receive Maximum slot across reconnects. `disconnect()` is no longer delayed while publishes wait for quota (`[MQTT-3.3.4-8]`). +- **Topic Alias Maximum is enforced and reset on each CONNACK** (`[MQTT-3.2.2-17]`, `[MQTT-3.2.2-18]`). Inbound Topic Aliases are resolved per connection (`[MQTT-3.3.2-10]`). An inbound alias of 0 or above the client maximum is a protocol error. Before, an aliased PUBLISH with an empty topic was acknowledged and then dropped. +- **Protocol errors close the connection properly.** On a malformed packet or protocol violation the client sends DISCONNECT with the matching reason code (0x81, 0x82, 0x93, 0x94, 0x95), flushes it, and closes the network connection. Previously it only marked itself disconnected, kept the socket open, and kept sending PINGREQ. A server DISCONNECT now closes the connection too (`[MQTT-4.13.2-1]`). Nothing is written after the client's own DISCONNECT (`[MQTT-3.14.4-1]`). +- **Inbound checks added:** + - reserved flags on SUBACK, PUBLISH, SUBSCRIBE and UNSUBSCRIBE (`[MQTT-2.1.3-1]`) + - the client's advertised Receive Maximum (0x93) and Maximum Packet Size (0x95) + - Subscription Identifier 0 + - Request Problem Information=0: a Reason String or User Property on a packet other than PUBLISH, CONNACK or DISCONNECT is a protocol error +- **Acknowledgements go out in PUBLISH arrival order when deferred ack is enabled** (`[MQTT-4.6.0-2]`, `[MQTT-4.6.0-3]`). Automatic acks and `AckToken` acks now share one ordered release, so a later ack waits for earlier pending ones. `AckToken::reject` maps reason codes that are invalid for PUBACK/PUBREC to 0x80. Acks still queued when the connection drops are discarded if the new connection reports Session Present=0. +- **WebSocket reads reassemble MQTT packets from the byte stream** (`[MQTT-6.0.0-2]`). Several packets in one frame, or one packet split across frames, used to corrupt payloads or drop the session. A text frame closes the connection (`[MQTT-6.0.0-1]`), and a WebSocket Ping no longer ends the session. +- **CONNECT carries `request_problem_information`, `request_response_information` and user properties**, which were silently dropped. The client never sends AUTH when CONNECT had no Authentication Method (`[MQTT-4.12.0-7]`). The Assigned Client Identifier is adopted for later reconnects (`[MQTT-3.1.3-2]`). + +## [mqttv5-cli 0.28.8] - 2026-09-23 + +### Changed + +- Depends on `mqtt5` 0.41. `pub` and `sub` with `--no-clean-start` set `resume_existing_session`, so they resume the broker-held session as before. + +## [mqtt5-wasm 1.5.0] - 2026-09-23 + +### Added + +- **`ConnectOptions.resumeExistingSession`.** It lets a freshly created client accept Session Present=1 from a broker-held session. The `session-recovery` and `qos2-recovery` examples use it. + +### Changed + +- **The browser client now follows the same client-side conformance rules as the native client.** This release fixes the missing PUBACK for inbound QoS 1 messages, byte loss when several packets arrived in one frame, packet identifier reuse, and the missing resend on session resume. It also adds topic and filter validation, and enforcement of the server's Receive Maximum, Topic Alias Maximum, Maximum QoS, Retain Available and Maximum Packet Size. Protocol errors now send DISCONNECT with a reason code and close the transport. With `keepAlive=0`, the client no longer sends PINGREQs. A fresh client now rejects Session Present=1 unless `resumeExistingSession` is set. Invalid requests are rejected before sending, and QoS 1/2 publishes wait while the server's Receive Maximum is exhausted. Tests: `crates/mqtt5-wasm/tests/conformance_client.rs`. +- Depends on `mqtt5` 0.41 and `mqtt5-protocol` 0.15.2. + +## [mqtt5-protocol 0.15.2] - 2026-09-23 + +### Added + +- **`PacketIdGenerator::next_available(in_use)`** returns the next identifier that isn't in use, or `None` when all are taken. +- **`validation::is_valid_subscription_filter` / `validate_subscription_filter`** validate topic filters, including the `$share/{ShareName}/{filter}` rules (`[MQTT-4.8.2-1]`, `[MQTT-4.8.2-2]`). + +### Fixed + +- Packet decoding now checks fixed-header reserved flags for every packet type (`[MQTT-2.1.3-1]`). +- `TopicAliasManager` no longer overflows when the alias maximum is 65535. + ## [mqtt5 0.40.0] - 2026-09-08 ### Breaking diff --git a/crates/mqtt5-conformance/CONFORMANCE_DIARY.md b/crates/mqtt5-conformance/CONFORMANCE_DIARY.md index 4ae5fc37..94b649c4 100644 --- a/crates/mqtt5-conformance/CONFORMANCE_DIARY.md +++ b/crates/mqtt5-conformance/CONFORMANCE_DIARY.md @@ -38,6 +38,20 @@ ## Diary Entries +### Client-side audit: the suite only ever tested brokers, and our own clients failed ~40 MUSTs (2026-09-23) + +**Trigger**: checking a third-party client's claim of full MQTT v5 conformance. This suite's SUT is always a broker, so it could not answer the question. The 149 statements in `conformance.toml` with `applies_to = "Client"` or `"Both"` were audited instead with raw-byte fake-broker tests that drive the real client and record what it puts on the wire. After the third-party client had been tested, the same audit ran against our own `MqttClient` and the `mqtt5-wasm` client. + +**Result**: our native client failed about 40 client MUST statements. The third-party client failed 7. Resend on session resume did not exist. Packet identifiers were reused while in flight. There was no topic or filter validation. Topic Alias Maximum and Retain Available were never enforced. The offline queue bypassed flow control. Protocol errors left the socket half-open. WebSocket reads assumed one packet per frame. The wasm client had most of the same defects and also never sent PUBACK. All of them are fixed in mqtt5 0.41.0 and mqtt5-wasm 1.5.0. + +**Where the tests live**: `crates/mqtt5/tests/conf_client_{a,b,c,d}.rs` (native) and `crates/mqtt5-wasm/tests/conformance_client.rs` (wasm, MessagePort fake broker under Node). They are not yet registered in this crate's manifest or runner. A client-side SUT mode for this suite is the natural next step. + +**Manifest drift bit the audit**: statement IDs in `conformance.toml` were used to label findings, and several were wrong. For example, the manifest files "no session state + Session Present=1 → close" under 3.2.2-5, but it is 3.2.2-4. All client-test names use IDs from `mqtt-v5.0-statement-texts.txt`. `known-text-drift.txt` is real debt with consequences outside this crate. + +**Decision recorded**: `[MQTT-3.2.2-4]` is enforced strictly by default. A fresh client can still resume a broker-held session through an explicit `ConnectOptions::resume_existing_session` opt-in, which the deferred-ack crash-recovery pattern needs. This crate's in-process test client sets the opt-in, because it checks `session_present` as an observer of the broker. + +**Lesson**: a conformance suite that only tests one side of the protocol says nothing about the other. Run any check we would point at someone else's implementation against our own first. + ### External-broker ack timeouts were a lost-wakeup race in the test client, not broker timing (2026-09-07) **Trigger**: issue #146. `deferred_qos2_zero_quota_still_serves_control_plane [MQTT-4.9.0-3]` failed diff --git a/crates/mqtt5-conformance/src/test_client/inprocess.rs b/crates/mqtt5-conformance/src/test_client/inprocess.rs index 503eba21..80583aff 100644 --- a/crates/mqtt5-conformance/src/test_client/inprocess.rs +++ b/crates/mqtt5-conformance/src/test_client/inprocess.rs @@ -7,7 +7,7 @@ use super::{MessageQueue, ReceivedMessage, Subscription, TestClientError}; use crate::sut::SutHandle; use mqtt5::MqttClient; use mqtt5_protocol::types::{ConnectOptions, PublishOptions, SubscribeOptions}; -use std::sync::{Arc, Mutex}; +use std::sync::{Arc, Mutex, PoisonError}; /// In-process backing for [`crate::test_client::TestClient`]. pub struct InProcessTestClient { @@ -35,6 +35,7 @@ impl InProcessTestClient { let wrapper = mqtt5::ConnectOptions { protocol_options: options.clone(), + resume_existing_session: true, ..mqtt5::ConnectOptions::default() }; let client = MqttClient::with_options(wrapper.clone()); @@ -79,10 +80,6 @@ impl InProcessTestClient { /// /// # Errors /// Returns an error if the broker rejects the subscription. - /// - /// # Panics - /// Panics from the delivery callback if the internal mutex has been - /// poisoned. pub async fn subscribe( &self, filter: &str, @@ -95,7 +92,7 @@ impl InProcessTestClient { .subscribe_with_options(filter, options, move |msg| { messages_cb .lock() - .unwrap() + .unwrap_or_else(PoisonError::into_inner) .push(ReceivedMessage::from_message(msg)); }) .await?; diff --git a/crates/mqtt5-protocol/Cargo.toml b/crates/mqtt5-protocol/Cargo.toml index 5050b7fd..59eb17a1 100644 --- a/crates/mqtt5-protocol/Cargo.toml +++ b/crates/mqtt5-protocol/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "mqtt5-protocol" -version = "0.15.1" +version = "0.15.2" edition.workspace = true rust-version.workspace = true authors.workspace = true diff --git a/crates/mqtt5-protocol/src/packet.rs b/crates/mqtt5-protocol/src/packet.rs index 9a1fa9e3..51ff9f92 100644 --- a/crates/mqtt5-protocol/src/packet.rs +++ b/crates/mqtt5-protocol/src/packet.rs @@ -193,7 +193,6 @@ impl PacketType { /// Converts a u8 to `PacketType` #[must_use] pub fn from_u8(value: u8) -> Option { - // Use the TryFrom implementation generated by BeBytes Self::try_from(value).ok() } } @@ -271,9 +270,9 @@ impl FixedHeader { #[must_use] pub fn validate_flags(&self) -> bool { match self.packet_type { - PacketType::Publish => true, // Publish has variable flags + PacketType::Publish => true, PacketType::PubRel | PacketType::Subscribe | PacketType::Unsubscribe => { - self.flags == 0x02 // Required flags for these packet types + self.flags == 0x02 } _ => self.flags == 0, } @@ -282,7 +281,6 @@ impl FixedHeader { /// Returns the encoded length of the fixed header #[must_use] pub fn encoded_len(&self) -> usize { - // 1 byte for packet type + flags, plus variable length encoding of remaining length 1 + crate::encoding::encoded_variable_int_len(self.remaining_length) } } @@ -413,6 +411,13 @@ impl Packet { buf: &mut B, protocol_version: u8, ) -> Result { + if !fixed_header.validate_flags() { + return Err(MqttError::MalformedPacket(format!( + "Invalid fixed header flags 0x{:02X} for {:?}", + fixed_header.flags, packet_type + ))); + } + match packet_type { PacketType::Publish => { let packet = publish::PublishPacket::decode_body_with_version( @@ -481,7 +486,6 @@ pub trait MqttPacket: Sized { /// /// Returns an error if encoding fails fn encode(&self, buf: &mut B) -> Result<()> { - // First encode to temporary buffer to get remaining length let mut body = Vec::new(); self.encode_body(&mut body)?; @@ -555,6 +559,25 @@ mod tests { assert!(header.validate_flags()); } + #[test] + fn test_versioned_decode_rejects_reserved_flags() { + let cases = [ + (PacketType::SubAck, 0x02, vec![0x00, 0x01, 0x00, 0x00]), + (PacketType::UnsubAck, 0x01, vec![0x00, 0x01, 0x00, 0x00]), + (PacketType::Subscribe, 0x00, vec![0x00, 0x01, 0x00]), + (PacketType::Unsubscribe, 0x00, vec![0x00, 0x01, 0x00]), + ]; + for (packet_type, flags, body) in cases { + let header = FixedHeader::new(packet_type, flags, u32::try_from(body.len()).unwrap()); + let mut buf = BytesMut::from(&body[..]); + let result = Packet::decode_from_body_with_version(packet_type, &header, &mut buf, 5); + assert!( + matches!(result, Err(MqttError::MalformedPacket(_))), + "{packet_type:?} with flags 0x{flags:02X} must be malformed: {result:?}" + ); + } + } + #[test] fn test_decode_insufficient_data() { let mut buf = BytesMut::new(); @@ -565,8 +588,8 @@ mod tests { #[test] fn test_decode_invalid_packet_type() { let mut buf = BytesMut::new(); - buf.put_u8(0x00); // Invalid packet type 0 - buf.put_u8(0x00); // Remaining length + buf.put_u8(0x00); + buf.put_u8(0x00); let result = FixedHeader::decode(&mut buf); assert!(result.is_err()); @@ -574,7 +597,6 @@ mod tests { #[test] fn test_packet_type_bebytes_serialization() { - // Test BeBytes to_be_bytes and try_from_be_bytes let packet_type = PacketType::Publish; let bytes = packet_type.to_be_bytes(); assert_eq!(bytes, vec![3]); @@ -583,7 +605,6 @@ mod tests { assert_eq!(decoded, PacketType::Publish); assert_eq!(consumed, 1); - // Test other packet types let packet_type = PacketType::Connect; let bytes = packet_type.to_be_bytes(); assert_eq!(bytes, vec![1]); diff --git a/crates/mqtt5-protocol/src/packet_id.rs b/crates/mqtt5-protocol/src/packet_id.rs index ce18ce2c..c8cd13db 100644 --- a/crates/mqtt5-protocol/src/packet_id.rs +++ b/crates/mqtt5-protocol/src/packet_id.rs @@ -52,6 +52,11 @@ impl PacketIdGenerator { current } } + + #[must_use] + pub fn next_available(&self, in_use: impl Fn(u16) -> bool) -> Option { + (0..u16::MAX).map(|_| self.next()).find(|id| !in_use(*id)) + } } impl Default for PacketIdGenerator { @@ -83,6 +88,26 @@ mod tests { assert_eq!(gen.next(), 2); } + #[test] + fn test_next_available_skips_ids_in_use() { + let gen = PacketIdGenerator::new(); + assert_eq!(gen.next_available(|id| id == 1 || id == 2), Some(3)); + assert_eq!(gen.next_available(|_| false), Some(4)); + } + + #[test] + fn test_next_available_skips_in_use_id_after_wrap() { + let gen = PacketIdGenerator::new(); + gen.next_id.store(u16::MAX, Ordering::SeqCst); + assert_eq!(gen.next_available(|id| id == u16::MAX || id == 1), Some(2)); + } + + #[test] + fn test_next_available_exhausted() { + let gen = PacketIdGenerator::new(); + assert_eq!(gen.next_available(|_| true), None); + } + #[cfg(feature = "std")] #[test] fn test_concurrent_access() { diff --git a/crates/mqtt5-protocol/src/session/topic_alias.rs b/crates/mqtt5-protocol/src/session/topic_alias.rs index c0a3b401..7095872c 100644 --- a/crates/mqtt5-protocol/src/session/topic_alias.rs +++ b/crates/mqtt5-protocol/src/session/topic_alias.rs @@ -32,27 +32,26 @@ impl TopicAliasManager { return None; } - while self.alias_to_topic.contains_key(&self.next_alias) - && self.next_alias <= self.topic_alias_maximum - { - self.next_alias += 1; - if self.next_alias > self.topic_alias_maximum { - self.next_alias = 1; - } + while self.alias_to_topic.contains_key(&self.next_alias) { + self.next_alias = self.alias_after(self.next_alias); } let alias = self.next_alias; self.alias_to_topic.insert(alias, topic.to_string()); self.topic_to_alias.insert(topic.to_string(), alias); - - self.next_alias += 1; - if self.next_alias > self.topic_alias_maximum { - self.next_alias = 1; - } + self.next_alias = self.alias_after(alias); Some(alias) } + fn alias_after(&self, alias: u16) -> u16 { + if alias >= self.topic_alias_maximum { + 1 + } else { + alias + 1 + } + } + /// # Errors /// Returns `TopicAliasInvalid` if alias is 0 or exceeds the maximum. pub fn register_alias(&mut self, alias: u16, topic: &str) -> Result<()> { @@ -141,6 +140,24 @@ mod tests { assert!(alias3.is_none()); } + #[test] + fn test_topic_alias_assigns_maximum_u16_without_overflow() { + let mut ta = TopicAliasManager::new(u16::MAX); + let mut last = None; + for i in 0..u32::from(u16::MAX) { + last = ta.get_or_create_alias(&format!("t/{i}")); + } + assert_eq!(last, Some(u16::MAX)); + assert_eq!(ta.get_or_create_alias("t/overflow"), None); + } + + #[test] + fn test_topic_alias_skips_registered_alias() { + let mut ta = TopicAliasManager::new(2); + ta.register_alias(1, "t/1").unwrap(); + assert_eq!(ta.get_or_create_alias("t/2"), Some(2)); + } + #[test] fn test_topic_alias_clear() { let mut ta = TopicAliasManager::new(10); diff --git a/crates/mqtt5-protocol/src/validation/mod.rs b/crates/mqtt5-protocol/src/validation/mod.rs index bcb24189..2179b6a3 100644 --- a/crates/mqtt5-protocol/src/validation/mod.rs +++ b/crates/mqtt5-protocol/src/validation/mod.rs @@ -28,7 +28,6 @@ pub fn is_valid_topic_name(topic: &str) -> bool { return false; } - // Topic names should not contain wildcards if topic.contains('+') || topic.contains('#') { return false; } @@ -60,30 +59,52 @@ pub fn is_valid_topic_filter(filter: &str) -> bool { let parts: Vec<&str> = filter.split('/').collect(); for (i, part) in parts.iter().enumerate() { - // Multi-level wildcard rules if part.contains('#') { - // # must be the last character in the filter if i != parts.len() - 1 { return false; } - // # must occupy the entire level if *part != "#" { return false; } } - // Single-level wildcard rules - if part.contains('+') { - // + must occupy the entire level - if *part != "+" { - return false; - } + if part.contains('+') && *part != "+" { + return false; } } true } +#[must_use] +pub fn is_valid_subscription_filter(filter: &str) -> bool { + let first_level = filter.split('/').next().unwrap_or(filter); + if first_level != "$share" { + return is_valid_topic_filter(filter); + } + if filter.len() > crate::constants::limits::MAX_STRING_LENGTH as usize { + return false; + } + filter + .strip_prefix("$share/") + .and_then(|rest| rest.split_once('/')) + .is_some_and(|(share_name, topic_filter)| { + !share_name.is_empty() + && !share_name.contains(['+', '#']) + && is_valid_topic_filter(topic_filter) + }) +} + +/// # Errors +/// +/// Returns `MqttError::InvalidTopicFilter` if [`is_valid_subscription_filter`] rejects it +pub fn validate_subscription_filter(filter: &str) -> Result<()> { + if !is_valid_subscription_filter(filter) { + return Err(MqttError::InvalidTopicFilter(filter.to_string())); + } + Ok(()) +} + /// Validates an MQTT client identifier according to MQTT v5.0 specification /// /// # Rules: @@ -94,18 +115,13 @@ pub fn is_valid_topic_filter(filter: &str) -> bool { #[must_use] pub fn is_valid_client_id(client_id: &str) -> bool { if client_id.is_empty() { - return true; // Empty client ID is allowed + return true; } - if client_id.len() > 23 { - // Most servers support longer, but 23 is the spec minimum - // We'll allow longer and let the server reject if needed - if client_id.len() > crate::constants::limits::MAX_CLIENT_ID_LENGTH { - return false; // Reasonable upper limit - } + if client_id.len() > crate::constants::limits::MAX_CLIENT_ID_LENGTH { + return false; } - // Check for valid characters (alphanumeric) client_id.chars().all(|c| c.is_ascii_alphanumeric()) } @@ -190,7 +206,6 @@ pub fn validate_client_id(client_id: &str) -> Result<()> { /// - Topics starting with '$' do NOT match root-level wildcards (MQTT spec) #[must_use] pub fn topic_matches_filter(topic: &str, filter: &str) -> bool { - // MQTT spec: topics starting with $ do not match wildcards at root level if topic.starts_with('$') && (filter.starts_with('#') || filter.starts_with('+')) { return false; } @@ -207,23 +222,21 @@ pub fn topic_matches_filter(topic: &str, filter: &str) -> bool { while t_idx < topic_parts.len() && f_idx < filter_parts.len() { if filter_parts[f_idx] == "#" { - return true; // Multi-level wildcard matches everything remaining + return true; } if filter_parts[f_idx] != "+" && filter_parts[f_idx] != topic_parts[t_idx] { - return false; // Not a match + return false; } t_idx += 1; f_idx += 1; } - // Check if we've consumed both topic and filter if t_idx == topic_parts.len() && f_idx == filter_parts.len() { return true; } - // Check if filter ends with # and we've consumed the topic if t_idx == topic_parts.len() && f_idx == filter_parts.len() - 1 && filter_parts[f_idx] == "#" { return true; } @@ -289,7 +302,6 @@ impl TopicValidator for StandardValidator { } fn is_reserved_topic(&self, _topic: &str) -> bool { - // Standard MQTT has no reserved topics false } @@ -351,7 +363,6 @@ impl RestrictiveValidator { /// Checks if topic violates additional restrictions fn check_additional_restrictions(&self, topic: &str) -> Result<()> { - // Check reserved prefixes for prefix in &self.reserved_prefixes { if topic.starts_with(prefix) { return Err(MqttError::InvalidTopicName(format!( @@ -360,7 +371,6 @@ impl RestrictiveValidator { } } - // Check maximum levels if let Some(max_levels) = self.max_levels { let level_count = topic.split('/').count(); if level_count > max_levels { @@ -370,7 +380,6 @@ impl RestrictiveValidator { } } - // Check maximum length if let Some(max_length) = self.max_topic_length { if topic.len() > max_length { return Err(MqttError::InvalidTopicName(format!( @@ -382,7 +391,6 @@ impl RestrictiveValidator { } } - // Check prohibited characters for &prohibited_char in &self.prohibited_chars { if topic.contains(prohibited_char) { return Err(MqttError::InvalidTopicName(format!( @@ -397,19 +405,13 @@ impl RestrictiveValidator { impl TopicValidator for RestrictiveValidator { fn validate_topic_name(&self, topic: &str) -> Result<()> { - // First apply standard validation validate_topic_name(topic)?; - // Then apply additional restrictions self.check_additional_restrictions(topic) } fn validate_topic_filter(&self, filter: &str) -> Result<()> { - // First apply standard validation validate_topic_filter(filter)?; - // Then apply additional restrictions (but allow wildcards) - // Note: We don't apply all restrictions to filters since they may contain wildcards - // Check reserved prefixes for prefix in &self.reserved_prefixes { if filter.starts_with(prefix) && !filter.contains('+') && !filter.contains('#') { return Err(MqttError::InvalidTopicFilter(format!( @@ -477,6 +479,49 @@ mod tests { assert!(!is_valid_topic_filter("home\0temperature")); } + #[test] + fn test_valid_subscription_filters() { + for filter in [ + "$shared/x", + "$sharex", + "$share/g/a/+", + "$share/g/#", + "$share/g//", + "/", + "+/+", + "a/+/b/#", + "+", + "#", + "$SYS/#", + "a//b", + ] { + assert!(is_valid_subscription_filter(filter), "{filter}"); + } + } + + #[test] + fn test_invalid_subscription_filters() { + for filter in [ + "", + "a/#/b", + "a+", + "$share", + "$share/", + "$share//x", + "$share/g", + "$share/g/", + "$share/g/a/#/b", + "$share/g+/x", + "$share/g#/x", + "$share/+/x", + "$share/#/x", + "$share/g/a\0", + ] { + assert!(!is_valid_subscription_filter(filter), "{filter:?}"); + assert!(validate_subscription_filter(filter).is_err()); + } + } + #[test] fn test_valid_client_ids() { assert!(is_valid_client_id("")); @@ -526,10 +571,8 @@ mod tests { #[test] fn test_topic_matches_filter() { - // Exact matches assert!(topic_matches_filter("sport/tennis", "sport/tennis")); - // Single-level wildcard assert!(topic_matches_filter("sport/tennis", "sport/+")); assert!(topic_matches_filter( "sport/tennis/player1", @@ -541,7 +584,6 @@ mod tests { )); assert!(!topic_matches_filter("sport/tennis/player1", "sport/+")); - // Multi-level wildcard assert!(topic_matches_filter("sport/tennis", "sport/#")); assert!(topic_matches_filter("sport/tennis/player1", "sport/#")); assert!(topic_matches_filter( @@ -552,7 +594,6 @@ mod tests { assert!(topic_matches_filter("anything", "#")); assert!(topic_matches_filter("sport/tennis", "#")); - // $ prefix topics - MQTT spec compliant behavior assert!(!topic_matches_filter("$SYS/broker/uptime", "#")); assert!(!topic_matches_filter( "$SYS/broker/uptime", @@ -562,7 +603,6 @@ mod tests { assert!(topic_matches_filter("$SYS/broker/uptime", "$SYS/#")); assert!(topic_matches_filter("$SYS/broker/uptime", "$SYS/+/uptime")); - // Non-matches assert!(!topic_matches_filter("sport/tennis", "sport/football")); assert!(!topic_matches_filter("sport", "sport/tennis")); assert!(!topic_matches_filter( diff --git a/crates/mqtt5-wasm/Cargo.toml b/crates/mqtt5-wasm/Cargo.toml index 0d7918fb..1e5c935b 100644 --- a/crates/mqtt5-wasm/Cargo.toml +++ b/crates/mqtt5-wasm/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "mqtt5-wasm" -version = "1.4.6" +version = "1.5.0" edition.workspace = true rust-version.workspace = true authors.workspace = true @@ -27,8 +27,8 @@ broker = ["client", "dep:mqtt5", "mqtt5/broker", "dep:tokio"] codec = ["client", "dep:miniz_oxide"] [dependencies] -mqtt5-protocol = "0.15.1" -mqtt5 = { version = "0.40", optional = true, default-features = false, features = [ +mqtt5-protocol = "0.15.2" +mqtt5 = { version = "0.41", optional = true, default-features = false, features = [ "tokio", ] } @@ -54,5 +54,8 @@ features = ["BinaryType", "Blob", "BroadcastChannel", "CloseEvent", "ErrorEvent" [target.'cfg(target_arch = "wasm32")'.dependencies] tokio = { version = "1.47", optional = true, features = ["sync"] } +[dev-dependencies] +wasm-bindgen-test = "0.3.78" + [package.metadata.wasm-pack.profile.release] wasm-opt = false diff --git a/crates/mqtt5-wasm/examples/README.md b/crates/mqtt5-wasm/examples/README.md index d8e94818..2d941346 100644 --- a/crates/mqtt5-wasm/examples/README.md +++ b/crates/mqtt5-wasm/examples/README.md @@ -32,7 +32,7 @@ Demonstrates QoS 2 (exactly once) message delivery with full acknowledgment flow ### QoS 2 Recovery (qos2-recovery/) -Demonstrates QoS 2 mid-flight recovery after client disconnection. Subscriber disconnects during a QoS 2 exchange, then reconnects with cleanStart=false to receive the in-flight message from persistent session state. +Demonstrates QoS 2 mid-flight recovery after client disconnection. Subscriber disconnects during a QoS 2 exchange, then reconnects with a fresh client instance using cleanStart=false and resumeExistingSession=true to receive the in-flight message from persistent session state. ### Will Message (will-message/) @@ -56,7 +56,7 @@ Demonstrates bandwidth optimization using topic aliases to reduce packet size fo ### Session Recovery (session-recovery/) -Shows session persistence across reconnections with cleanStart=false and sessionExpiryInterval. +Shows session persistence across reconnections with cleanStart=false, resumeExistingSession and sessionExpiryInterval. ### $SYS Monitoring (sys-monitoring/) diff --git a/crates/mqtt5-wasm/examples/qos2-recovery/index.html b/crates/mqtt5-wasm/examples/qos2-recovery/index.html index e622b5f2..45453e0a 100644 --- a/crates/mqtt5-wasm/examples/qos2-recovery/index.html +++ b/crates/mqtt5-wasm/examples/qos2-recovery/index.html @@ -416,6 +416,7 @@

Published Messages:

const opts = new ConnectOptions(); opts.cleanStart = false; + opts.resumeExistingSession = true; opts.sessionExpiryInterval = 3600; subscriberClient.connectMessagePortWithOptions(subscriberPort, opts).then(() => { diff --git a/crates/mqtt5-wasm/examples/session-recovery/index.html b/crates/mqtt5-wasm/examples/session-recovery/index.html index 481a448e..d33f9570 100644 --- a/crates/mqtt5-wasm/examples/session-recovery/index.html +++ b/crates/mqtt5-wasm/examples/session-recovery/index.html @@ -204,6 +204,7 @@

MQTT Session Persistence & Recovery Demo

How Session Persistence Works:
  • cleanStart=false: Resume existing session, keep subscriptions and queued messages
  • +
  • resumeExistingSession=true: Lets a fresh client instance (e.g. after a page reload or crash) accept the session the broker still holds
  • cleanStart=true: Start fresh, discard any existing session state
  • sessionExpiryInterval: How long the broker keeps the session after disconnect
@@ -437,6 +438,7 @@

cleanStart=false

const options = new ConnectOptions(); options.cleanStart = cleanStart; + options.resumeExistingSession = !cleanStart; options.sessionExpiryInterval = 3600; await mobileClient.connectMessagePortWithOptions(mobilePort, options); diff --git a/crates/mqtt5-wasm/examples/websocket/app.js b/crates/mqtt5-wasm/examples/websocket/app.js index e282d247..cba06a4f 100644 --- a/crates/mqtt5-wasm/examples/websocket/app.js +++ b/crates/mqtt5-wasm/examples/websocket/app.js @@ -288,7 +288,7 @@ async function handleSubscribe(e) { subOpts.noLocal = false; subOpts.retainAsPublished = true; subOpts.retainHandling = 0; - subOpts.subscriptionIdentifier = Math.floor(Math.random() * 1000000); + subOpts.subscriptionIdentifier = Math.floor(Math.random() * 1000000) + 1; console.log('handleSubscribe: Calling subscribeWithOptions for topic:', topic); const packetId = await client.subscribeWithOptions(topic, (receivedTopic, payload) => { diff --git a/crates/mqtt5-wasm/src/client/callbacks.rs b/crates/mqtt5-wasm/src/client/callbacks.rs index 03a15cea..97aea9e5 100644 --- a/crates/mqtt5-wasm/src/client/callbacks.rs +++ b/crates/mqtt5-wasm/src/client/callbacks.rs @@ -1,46 +1,85 @@ +use mqtt5_protocol::packet::disconnect::DisconnectPacket; +use mqtt5_protocol::packet::Packet; use std::cell::RefCell; use std::rc::Rc; use wasm_bindgen::prelude::*; +use super::handlers::Violation; +use super::packet::write_packet; +use super::qos::wake_quota_waiters; use super::reconnect::spawn_reconnection_task; use super::state::ClientState; -pub fn drain_pending_callbacks(state: &mut ClientState) { - let pubacks: Vec<_> = state.pending_pubacks.drain().collect(); - let pubcomps: Vec<_> = state.pending_pubcomps.drain().collect(); - let subacks: Vec<_> = state.pending_subacks.drain().collect(); +const SESSION_DISCARDED_REASON: u8 = 0x80; - let error_val = JsValue::from_f64(f64::from(0x80_u8)); - for (_, callback) in pubacks { - let _ = callback.call1(&JsValue::NULL, &error_val); +fn reject_callbacks(callbacks: Vec) { + let error_val = JsValue::from_f64(f64::from(SESSION_DISCARDED_REASON)); + for callback in callbacks { + if let Err(e) = callback.call1(&JsValue::NULL, &error_val) { + tracing::warn!(error = ?e, "acknowledgement callback failed"); + } + } +} + +pub fn discard_session(state: &Rc>) { + let callbacks = state.borrow_mut().discard_session(); + reject_callbacks(callbacks); +} + +pub fn close_network_connection(state: &Rc>) { + let writer = { + let mut state_mut = state.borrow_mut(); + state_mut.connected = false; + state_mut.connection_generation = state_mut.connection_generation.wrapping_add(1); + state_mut.writer.take() + }; + if let Some(writer) = writer { + if let Err(e) = writer.borrow_mut().close() { + tracing::warn!(error = %e, "closing the network connection failed"); + } } - for (_, (callback, _)) in pubcomps { - let _ = callback.call1(&JsValue::NULL, &error_val); +} + +pub fn end_connection_state(state: &Rc>) { + let (subacks, session_ends) = { + let mut state_mut = state.borrow_mut(); + state_mut.pending_unsubacks.clear(); + let subacks: Vec = state_mut + .pending_subacks + .drain() + .filter_map(|(_, callback)| callback) + .collect(); + (subacks, state_mut.session_expiry_interval == 0) + }; + for resolve in subacks { + let codes = js_sys::Array::new(); + codes.push(&JsValue::from_f64(f64::from(SESSION_DISCARDED_REASON))); + if let Err(e) = resolve.call1(&JsValue::NULL, &codes.into()) { + tracing::warn!(error = ?e, "SUBACK callback failed"); + } } - for (_, resolve) in subacks { - let arr = js_sys::Array::new(); - arr.push(&JsValue::from_f64(f64::from(0x80_u8))); - let _ = resolve.call1(&JsValue::NULL, &arr.into()); + if session_ends { + discard_session(state); } + wake_quota_waiters(state); } pub fn handle_connection_lost(state: &Rc>, reason: &str) { let should_reconnect = { - let mut state_ref = state.borrow_mut(); + let state_ref = state.borrow(); if !state_ref.connected { return; } - state_ref.connected = false; - - drain_pending_callbacks(&mut state_ref); - state_ref.reconnect_config.enabled && !state_ref.user_initiated_disconnect && !state_ref.reconnecting && state_ref.last_url.is_some() }; - web_sys::console::error_1(&reason.into()); + close_network_connection(state); + end_connection_state(state); + + tracing::warn!(reason, "connection lost"); trigger_error_callback(state, reason); trigger_disconnect_callback(state); @@ -49,11 +88,28 @@ pub fn handle_connection_lost(state: &Rc>, reason: &str) { } } +pub fn fail_connection(state: &Rc>, violation: &Violation) { + if !state.borrow().connected { + return; + } + if state.borrow().protocol_version == 5 { + let disconnect = DisconnectPacket::new(violation.reason_code); + if let Err(e) = write_packet(state, &Packet::Disconnect(disconnect)) { + tracing::warn!(error = %e, "DISCONNECT not sent"); + } + } + let message = format!( + "Protocol violation ({:?}): {}", + violation.reason_code, violation.message + ); + handle_connection_lost(state, &message); +} + pub fn trigger_disconnect_callback(state: &Rc>) { let callback = state.borrow().on_disconnect.clone(); if let Some(callback) = callback { if let Err(e) = callback.call0(&JsValue::NULL) { - web_sys::console::error_1(&format!("onDisconnect callback error: {e:?}").into()); + tracing::warn!(error = ?e, "onDisconnect callback failed"); } } } @@ -63,7 +119,7 @@ pub fn trigger_error_callback(state: &Rc>, error_msg: &str) if let Some(callback) = callback { let error_js = JsValue::from_str(error_msg); if let Err(e) = callback.call1(&JsValue::NULL, &error_js) { - web_sys::console::error_1(&format!("onError callback error: {e:?}").into()); + tracing::warn!(error = ?e, "onError callback failed"); } } } @@ -78,7 +134,7 @@ pub fn trigger_reconnecting_callback( let attempt_js = JsValue::from_f64(f64::from(attempt)); let delay_js = JsValue::from_f64(f64::from(delay_millis)); if let Err(e) = callback.call2(&JsValue::NULL, &attempt_js, &delay_js) { - web_sys::console::error_1(&format!("onReconnecting callback error: {e:?}").into()); + tracing::warn!(error = ?e, "onReconnecting callback failed"); } } } @@ -88,9 +144,7 @@ pub fn trigger_connectivity_change_callback(state: &Rc>, on if let Some(callback) = callback { let online_js = JsValue::from_bool(online); if let Err(e) = callback.call1(&JsValue::NULL, &online_js) { - web_sys::console::error_1( - &format!("onConnectivityChange callback error: {e:?}").into(), - ); + tracing::warn!(error = ?e, "onConnectivityChange callback failed"); } } } @@ -100,7 +154,7 @@ pub fn trigger_reconnect_failed_callback(state: &Rc>, error if let Some(callback) = callback { let error_js = JsValue::from_str(error_msg); if let Err(e) = callback.call1(&JsValue::NULL, &error_js) { - web_sys::console::error_1(&format!("onReconnectFailed callback error: {e:?}").into()); + tracing::warn!(error = ?e, "onReconnectFailed callback failed"); } } } diff --git a/crates/mqtt5-wasm/src/client/connection.rs b/crates/mqtt5-wasm/src/client/connection.rs new file mode 100644 index 00000000..b016d161 --- /dev/null +++ b/crates/mqtt5-wasm/src/client/connection.rs @@ -0,0 +1,233 @@ +use bytes::BytesMut; +use mqtt5_protocol::packet::connack::ConnAckPacket; +use mqtt5_protocol::packet::connect::ConnectPacket; +use mqtt5_protocol::packet::disconnect::DisconnectPacket; +use mqtt5_protocol::packet::Packet; +use mqtt5_protocol::protocol::v5::reason_codes::ReasonCode; +use mqtt5_protocol::Transport; +use std::cell::RefCell; +use std::fmt; +use std::rc::Rc; +use wasm_bindgen::JsValue; + +use crate::transport::{WasmReader, WasmTransportType}; + +use super::callbacks::discard_session; +use super::handlers::handle_auth; +use super::keepalive::spawn_keepalive_task; +use super::packet::{encode_packet, write_packet}; +use super::qos::{resend_session, spawn_qos2_cleanup_task}; +use super::reader::{read_connect_response, spawn_packet_reader}; +use super::state::{ClientState, SessionState, StoredConnectOptions}; + +pub enum ConnectFailure { + Redirect(String), + Failed(String), +} + +impl fmt::Display for ConnectFailure { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Redirect(url) => write!(f, "Server redirected the client to {url}"), + Self::Failed(message) => f.write_str(message), + } + } +} + +impl From for JsValue { + fn from(failure: ConnectFailure) -> Self { + match failure { + ConnectFailure::Redirect(url) => { + let obj = js_sys::Object::new(); + js_sys::Reflect::set( + &obj, + &JsValue::from_str("type"), + &JsValue::from_str("redirect"), + ) + .ok(); + js_sys::Reflect::set(&obj, &JsValue::from_str("url"), &JsValue::from_str(&url)) + .ok(); + obj.into() + } + ConnectFailure::Failed(message) => JsValue::from_str(&message), + } + } +} + +pub async fn establish( + state: &Rc>, + mut transport: WasmTransportType, + connect: ConnectPacket, + options: &StoredConnectOptions, +) -> Result { + transport + .connect() + .await + .map_err(|e| ConnectFailure::Failed(format!("Transport connection failed: {e}")))?; + + if connect.clean_start { + discard_session(state); + } + state.borrow_mut().apply_connect_options(options); + + let mut buf = BytesMut::new(); + encode_packet(&Packet::Connect(Box::new(connect.clone())), &mut buf) + .map_err(|e| ConnectFailure::Failed(format!("Packet encoding failed: {e}")))?; + transport + .write(&buf) + .await + .map_err(|e| ConnectFailure::Failed(format!("Write failed: {e}")))?; + + let (mut reader, writer) = transport + .into_split() + .map_err(|e| ConnectFailure::Failed(format!("Transport split failed: {e}")))?; + state.borrow_mut().writer = Some(Rc::new(RefCell::new(writer))); + + match await_connack(state, &mut reader).await { + Ok(connack) => accept_connack( + state, + reader, + connack, + connect.clean_start, + options.resume_existing_session, + ), + Err(failure) => { + drop_writer(state); + Err(failure) + } + } +} + +fn drop_writer(state: &Rc>) { + let writer = state.borrow_mut().writer.take(); + if let Some(writer) = writer { + if let Err(e) = writer.borrow_mut().close() { + tracing::warn!(error = %e, "closing the network connection failed"); + } + } +} + +async fn await_connack( + state: &Rc>, + reader: &mut WasmReader, +) -> Result { + loop { + let packet = read_connect_response(state, reader) + .await + .map_err(|e| ConnectFailure::Failed(format!("Packet read failed: {e}")))?; + match packet { + Packet::ConnAck(connack) if connack.reason_code == ReasonCode::Success => { + return Ok(connack); + } + Packet::ConnAck(connack) => { + let reason_code = u8::from(connack.reason_code); + return Err( + match (reason_code, connack.properties.get_server_reference()) { + (0x9C | 0x9D, Some(server_reference)) => { + ConnectFailure::Redirect(server_reference.to_string()) + } + _ => ConnectFailure::Failed(format!( + "Connection rejected: {}", + connack_error_description(reason_code) + )), + }, + ); + } + Packet::Auth(auth) => handle_auth(state, &auth).map_err(|violation| { + ConnectFailure::Failed(format!( + "Authentication failed ({:?}): {}", + violation.reason_code, violation.message + )) + })?, + other => { + return Err(ConnectFailure::Failed(format!( + "Expected CONNACK or AUTH, received {}", + other.packet_type_name() + ))); + } + } + } +} + +fn accept_connack( + state: &Rc>, + reader: WasmReader, + connack: ConnAckPacket, + clean_start: bool, + resume_existing_session: bool, +) -> Result { + let (had_session, protocol_version) = { + let state_ref = state.borrow(); + ( + !clean_start && (resume_existing_session || state_ref.session == SessionState::Held), + state_ref.protocol_version, + ) + }; + + if connack.session_present && !had_session { + if protocol_version == 5 { + let disconnect = DisconnectPacket::new(ReasonCode::ProtocolError); + if let Err(e) = write_packet(state, &Packet::Disconnect(disconnect)) { + tracing::warn!(error = %e, "DISCONNECT not sent"); + } + } + drop_writer(state); + return Err(ConnectFailure::Failed( + "Server reported Session Present but the client holds no session state".to_string(), + )); + } + + if !connack.session_present { + discard_session(state); + } + + { + let mut state_mut = state.borrow_mut(); + state_mut.apply_connack(&connack); + state_mut.connected = true; + state_mut.session = SessionState::Held; + state_mut.connection_generation = state_mut.connection_generation.wrapping_add(1); + } + + if connack.session_present { + resend_session(state); + } + + spawn_packet_reader(Rc::clone(state), reader); + spawn_keepalive_task(Rc::clone(state)); + spawn_qos2_cleanup_task(Rc::clone(state)); + + let callback = state.borrow().on_connect.clone(); + if let Some(callback) = callback { + let reason_code_js = JsValue::from_f64(f64::from(u8::from(connack.reason_code))); + let session_present_js = JsValue::from_bool(connack.session_present); + if let Err(e) = callback.call2(&JsValue::NULL, &reason_code_js, &session_present_js) { + tracing::warn!(error = ?e, "onConnect callback failed"); + } + } + + Ok(connack) +} + +fn connack_error_description(reason_code: u8) -> &'static str { + match reason_code { + 0x80 => "Unspecified error", + 0x81 => "Malformed packet", + 0x82 => "Protocol error", + 0x83 => "Implementation specific error", + 0x84 => "Unsupported protocol version", + 0x85 => "Client identifier not valid", + 0x86 => "Bad username or password", + 0x87 => "Not authorized", + 0x88 => "Server unavailable", + 0x89 => "Server busy", + 0x8A => "Banned", + 0x8C => "Bad authentication method", + 0x90 => "Topic name invalid", + 0x97 => "Quota exceeded", + 0x9C => "Use another server", + 0x9D => "Server moved", + 0x9F => "Connection rate exceeded", + _ => "Unknown error", + } +} diff --git a/crates/mqtt5-wasm/src/client/handlers.rs b/crates/mqtt5-wasm/src/client/handlers.rs index 09ad23e5..fff07342 100644 --- a/crates/mqtt5-wasm/src/client/handlers.rs +++ b/crates/mqtt5-wasm/src/client/handlers.rs @@ -1,148 +1,286 @@ -use bytes::BytesMut; +use mqtt5_protocol::packet::puback::PubAckPacket; +use mqtt5_protocol::packet::pubcomp::PubCompPacket; +use mqtt5_protocol::packet::publish::PublishPacket; +use mqtt5_protocol::packet::pubrec::PubRecPacket; +use mqtt5_protocol::packet::pubrel::PubRelPacket; use mqtt5_protocol::packet::Packet; +use mqtt5_protocol::protocol::v5::properties::{Properties, PropertyId}; +use mqtt5_protocol::protocol::v5::reason_codes::ReasonCode; +use mqtt5_protocol::QoS; use std::cell::RefCell; use std::rc::Rc; use wasm_bindgen::prelude::*; -use wasm_bindgen_futures::spawn_local; use crate::config::WasmMessageProperties; -use super::packet::encode_packet; +use super::packet::write_packet; +use super::qos::release_quota; use super::state::ClientState; use super::RustMessage; #[cfg(feature = "codec")] use crate::codec::WasmCodecRegistry; -pub fn handle_incoming_packet(state: &Rc>, packet: Packet) { - match packet { - Packet::ConnAck(connack) => { - let callback = state.borrow().on_connect.clone(); - if let Some(callback) = callback { - let reason_code = JsValue::from_f64(f64::from(connack.reason_code as u8)); - let session_present = JsValue::from_bool(connack.session_present); +#[derive(Debug)] +pub struct Violation { + pub reason_code: ReasonCode, + pub message: String, +} - if let Err(e) = callback.call2(&JsValue::NULL, &reason_code, &session_present) { - web_sys::console::error_1(&format!("onConnect callback error: {e:?}").into()); - } - } +impl Violation { + pub fn new(reason_code: ReasonCode, message: impl Into) -> Self { + Self { + reason_code, + message: message.into(), } - Packet::Publish(ref publish) => handle_publish(state, publish), + } + + fn protocol(message: impl Into) -> Self { + Self::new(ReasonCode::ProtocolError, message) + } +} + +pub fn handle_incoming_packet( + state: &Rc>, + packet: Packet, +) -> Result<(), Violation> { + check_problem_information(state, &packet)?; + match packet { + Packet::Publish(publish) => handle_publish(state, &publish), Packet::SubAck(suback) => { let callback = state.borrow_mut().pending_subacks.remove(&suback.packet_id); - if let Some(callback) = callback { + if let Some(Some(callback)) = callback { let reason_codes = suback .reason_codes .iter() .map(|rc| JsValue::from_f64(f64::from(*rc as u8))) .collect::(); - if let Err(e) = callback.call1(&JsValue::NULL, &reason_codes.into()) { - web_sys::console::error_1(&format!("SUBACK callback error: {e:?}").into()); + tracing::warn!(error = ?e, "SUBACK callback failed"); } } + Ok(()) + } + Packet::UnsubAck(unsuback) => { + state + .borrow_mut() + .pending_unsubacks + .remove(&unsuback.packet_id); + Ok(()) } - Packet::UnsubAck(_unsuback) => {} Packet::PingResp => { state.borrow_mut().last_pong_received = Some(js_sys::Date::now()); + Ok(()) } Packet::PubAck(puback) => { - let callback = state.borrow_mut().pending_pubacks.remove(&puback.packet_id); - if let Some(callback) = callback { - let reason_code = JsValue::from_f64(f64::from(puback.reason_code as u8)); - if let Err(e) = callback.call1(&JsValue::NULL, &reason_code) { - web_sys::console::error_1(&format!("PUBACK callback error: {e:?}").into()); - } - } + handle_puback(state, puback.packet_id, puback.reason_code); + Ok(()) } Packet::PubRec(pubrec) => { handle_pubrec(state, pubrec.packet_id, pubrec.reason_code); + Ok(()) } Packet::PubComp(pubcomp) => { handle_pubcomp(state, pubcomp.packet_id, pubcomp.reason_code); + Ok(()) } Packet::PubRel(pubrel) => { handle_pubrel(state, pubrel.packet_id); + Ok(()) } - Packet::Auth(auth) => { - let reason_code = auth.reason_code; + Packet::Auth(auth) => handle_auth(state, &auth), + other => Err(Violation::protocol(format!( + "A client must not receive {}", + other.packet_type_name() + ))), + } +} - if reason_code - == mqtt5_protocol::protocol::v5::reason_codes::ReasonCode::ContinueAuthentication - { - let callback = state.borrow().on_auth_challenge.clone(); - if let Some(callback) = callback { - let auth_method = auth - .properties - .get_authentication_method() - .cloned() - .unwrap_or_default(); - let auth_data = auth.properties.get_authentication_data(); +fn check_problem_information( + state: &Rc>, + packet: &Packet, +) -> Result<(), Violation> { + if state.borrow().client_limits.request_problem_information { + return Ok(()); + } + let properties = match packet { + Packet::PubAck(p) => &p.properties, + Packet::PubRec(p) => &p.properties, + Packet::PubRel(p) => &p.properties, + Packet::PubComp(p) => &p.properties, + Packet::SubAck(p) => &p.properties, + Packet::UnsubAck(p) => &p.properties, + Packet::Auth(p) => &p.properties, + _ => return Ok(()), + }; + if carries_problem_information(properties) { + return Err(Violation::protocol(format!( + "{} carries a Reason String or User Property although Request Problem Information is 0", + packet.packet_type_name() + ))); + } + Ok(()) +} - let method_js = JsValue::from_str(&auth_method); - let data_js = if let Some(data) = auth_data { - js_sys::Uint8Array::from(data).into() - } else { - JsValue::NULL - }; +fn carries_problem_information(properties: &Properties) -> bool { + properties.contains(PropertyId::ReasonString) || properties.contains(PropertyId::UserProperty) +} - if let Err(e) = callback.call2(&JsValue::NULL, &method_js, &data_js) { - web_sys::console::error_1( - &format!("onAuthChallenge callback error: {e:?}").into(), - ); - } - } - } else if reason_code == mqtt5_protocol::protocol::v5::reason_codes::ReasonCode::Success +pub fn handle_auth( + state: &Rc>, + auth: &mqtt5_protocol::packet::auth::AuthPacket, +) -> Result<(), Violation> { + let (auth_method, callback) = { + let state_ref = state.borrow(); + ( + state_ref.auth_method.clone(), + state_ref.on_auth_challenge.clone(), + ) + }; + let Some(auth_method) = auth_method else { + return Err(Violation::protocol( + "AUTH received but CONNECT carried no Authentication Method", + )); + }; + if auth.properties.get_authentication_method() != Some(&auth_method) { + return Err(Violation::protocol( + "AUTH Authentication Method differs from CONNECT", + )); + } + match auth.reason_code { + ReasonCode::ContinueAuthentication => { + let Some(callback) = callback else { + return Err(Violation::new( + ReasonCode::ImplementationSpecificError, + "AUTH challenge received but no onAuthChallenge callback set", + )); + }; + let data_js = auth + .properties + .get_authentication_data() + .map_or(JsValue::NULL, |data| js_sys::Uint8Array::from(data).into()); + if let Err(e) = + callback.call2(&JsValue::NULL, &JsValue::from_str(&auth_method), &data_js) { - web_sys::console::log_1(&"Authentication successful".into()); + tracing::warn!(error = ?e, "onAuthChallenge callback failed"); } + Ok(()) } - _ => { - web_sys::console::warn_1(&format!("Unhandled packet type: {packet:?}").into()); + ReasonCode::Success => { + tracing::debug!("re-authentication succeeded"); + Ok(()) } + other => Err(Violation::protocol(format!( + "Unexpected AUTH reason code {other:?}" + ))), + } +} + +fn resolve_topic( + state: &Rc>, + publish: &PublishPacket, +) -> Result { + let mut state_mut = state.borrow_mut(); + match publish.topic_alias() { + None if publish.topic_name.is_empty() => Err(Violation::protocol( + "PUBLISH without Topic Name or Topic Alias", + )), + None => Ok(publish.topic_name.clone()), + Some(alias) if alias == 0 || alias > state_mut.client_limits.topic_alias_maximum => { + Err(Violation::new( + ReasonCode::TopicAliasInvalid, + format!( + "Topic Alias {alias} outside 1..={}", + state_mut.client_limits.topic_alias_maximum + ), + )) + } + Some(alias) if publish.topic_name.is_empty() => state_mut + .inbound_aliases + .get_topic(alias) + .map(str::to_string) + .ok_or_else(|| Violation::protocol(format!("Topic Alias {alias} has no mapping"))), + Some(alias) => { + state_mut + .inbound_aliases + .register_alias(alias, &publish.topic_name) + .map_err(|e| Violation::new(ReasonCode::TopicAliasInvalid, e.to_string()))?; + Ok(publish.topic_name.clone()) + } + } +} + +fn send_ack(state: &Rc>, packet: &Packet) { + if let Err(e) = write_packet(state, packet) { + tracing::warn!(error = %e, packet = packet.packet_type_name(), "acknowledgement not sent"); } } fn handle_publish( state: &Rc>, - publish: &mqtt5_protocol::packet::publish::PublishPacket, -) { - let topic = publish.topic_name.clone(); - let payload = publish.payload.clone(); - let qos = publish.qos; - let retain = publish.retain; + publish: &PublishPacket, +) -> Result<(), Violation> { + let topic = resolve_topic(state, publish)?; let properties: mqtt5_protocol::types::MessageProperties = publish.properties.clone().into(); - if qos == mqtt5_protocol::QoS::ExactlyOnce { - if let Some(packet_id) = publish.packet_id { - let is_duplicate = state.borrow().received_qos2.contains_key(&packet_id); - let actions = - mqtt5_protocol::qos2::handle_incoming_publish_qos2(packet_id, is_duplicate); - - for action in actions { - match action { - mqtt5_protocol::qos2::QoS2Action::DeliverMessage { packet_id: _ } => { - deliver_message(state, &topic, &payload, qos, retain, &properties); - } - mqtt5_protocol::qos2::QoS2Action::SendPubRec { - packet_id, - reason_code, - } => { - send_pubrec(state, packet_id, reason_code); - } - mqtt5_protocol::qos2::QoS2Action::TrackIncomingPubRec { packet_id } => { - let now = js_sys::Date::now(); - state.borrow_mut().pending_pubrecs.insert(packet_id, now); - state.borrow_mut().received_qos2.insert(packet_id, now); - } - _ => {} + match (publish.qos, publish.packet_id) { + (QoS::AtMostOnce, _) => { + deliver_message( + state, + &topic, + &publish.payload, + publish.qos, + publish.retain, + &properties, + ); + Ok(()) + } + (QoS::AtLeastOnce, Some(packet_id)) => { + deliver_message( + state, + &topic, + &publish.payload, + publish.qos, + publish.retain, + &properties, + ); + send_ack(state, &Packet::PubAck(PubAckPacket::new(packet_id))); + Ok(()) + } + (QoS::ExactlyOnce, Some(packet_id)) => { + let is_new = { + let mut state_mut = state.borrow_mut(); + if state_mut.awaiting_pubrel.contains(&packet_id) { + false + } else if state_mut.awaiting_pubrel.len() + >= usize::from(state_mut.client_limits.receive_maximum) + { + return Err(Violation::new( + ReasonCode::ReceiveMaximumExceeded, + "Server exceeded the client Receive Maximum", + )); + } else { + state_mut.awaiting_pubrel.insert(packet_id); + true } + }; + if is_new { + deliver_message( + state, + &topic, + &publish.payload, + publish.qos, + publish.retain, + &properties, + ); } - } else { - web_sys::console::error_1(&"QoS 2 PUBLISH missing packet_id".into()); + send_ack(state, &Packet::PubRec(PubRecPacket::new(packet_id))); + Ok(()) } - } else { - deliver_message(state, &topic, &payload, qos, retain, &properties); + (_, None) => Err(Violation::new( + ReasonCode::MalformedPacket, + "QoS > 0 PUBLISH without Packet Identifier", + )), } } @@ -150,7 +288,7 @@ fn deliver_message( state: &Rc>, topic: &str, payload: &[u8], - qos: mqtt5_protocol::QoS, + qos: QoS, retain: bool, properties: &mqtt5_protocol::types::MessageProperties, ) { @@ -182,7 +320,7 @@ fn deliver_message( &payload_array.into(), &props_js.into(), ) { - web_sys::console::error_1(&format!("Callback error: {e:?}").into()); + tracing::warn!(error = ?e, "message callback failed"); } } } @@ -215,152 +353,107 @@ fn decode_payload_if_needed( } } -fn send_pubrec( +fn complete_flight( state: &Rc>, packet_id: u16, - reason_code: mqtt5_protocol::protocol::v5::reason_codes::ReasonCode, + reason_code: ReasonCode, + callback: Option, ) { - let pubrec = - mqtt5_protocol::packet::pubrec::PubRecPacket::new_with_reason(packet_id, reason_code); - let mut buf = BytesMut::new(); - if let Err(e) = encode_packet(&mqtt5_protocol::packet::Packet::PubRec(pubrec), &mut buf) { - web_sys::console::error_1(&format!("PUBREC encode error: {e}").into()); - return; - } - - let writer_rc = state.borrow().writer.clone(); - if let Some(writer_rc) = writer_rc { - spawn_local(async move { - if let Err(e) = writer_rc.borrow_mut().write(&buf) { - web_sys::console::error_1(&format!("PUBREC send error: {e}").into()); - } - }); + release_quota(state); + if let Some(callback) = callback { + let reason_code_js = JsValue::from_f64(f64::from(u8::from(reason_code))); + if let Err(e) = callback.call1(&JsValue::NULL, &reason_code_js) { + tracing::warn!(error = ?e, packet_id, "publish acknowledgement callback failed"); + } } } -fn handle_pubrec( - state: &Rc>, - packet_id: u16, - reason_code: mqtt5_protocol::protocol::v5::reason_codes::ReasonCode, -) { - let has_pending = state.borrow().pending_pubcomps.contains_key(&packet_id); - let actions = mqtt5_protocol::qos2::handle_incoming_pubrec(packet_id, reason_code, has_pending); +fn handle_puback(state: &Rc>, packet_id: u16, reason_code: ReasonCode) { + let callback = { + let mut state_mut = state.borrow_mut(); + let acknowledged = state_mut + .outbound + .get(&packet_id) + .is_some_and(|flight| flight.publish.qos == QoS::AtLeastOnce); + if !acknowledged { + tracing::debug!(packet_id, "PUBACK for unknown packet identifier"); + return; + } + state_mut.outbound.remove(&packet_id); + state_mut.pending_pubacks.remove(&packet_id) + }; + complete_flight(state, packet_id, reason_code, callback); +} - for action in actions { - match action { - mqtt5_protocol::qos2::QoS2Action::SendPubRel { packet_id } => { - let pubrel_packet = mqtt5_protocol::packet::pubrel::PubRelPacket::new(packet_id); - let mut buf = BytesMut::new(); - if let Err(e) = encode_packet( - &mqtt5_protocol::packet::Packet::PubRel(pubrel_packet), - &mut buf, - ) { - web_sys::console::error_1(&format!("PUBREL encode error: {e}").into()); - continue; - } +enum PubRecOutcome { + Release, + Failed(Option), + Unknown, +} - let writer_rc = state.borrow().writer.clone(); - if let Some(writer_rc) = writer_rc { - spawn_local(async move { - if let Err(e) = writer_rc.borrow_mut().write(&buf) { - web_sys::console::error_1(&format!("PUBREL send error: {e}").into()); - } - }); - } +fn handle_pubrec(state: &Rc>, packet_id: u16, reason_code: ReasonCode) { + let outcome = { + let mut state_mut = state.borrow_mut(); + let failed = u8::from(reason_code) >= 0x80; + match state_mut.outbound.get_mut(&packet_id) { + Some(flight) if flight.publish.qos == QoS::ExactlyOnce && !failed => { + flight.released = true; + PubRecOutcome::Release } - mqtt5_protocol::qos2::QoS2Action::ErrorFlow { - packet_id, - reason_code, - } => { - if let Some((callback, _)) = state.borrow_mut().pending_pubcomps.remove(&packet_id) - { - let reason_code_js = JsValue::from_f64(f64::from(reason_code as u8)); - if let Err(e) = callback.call1(&JsValue::NULL, &reason_code_js) { - web_sys::console::error_1( - &format!("QoS 2 error callback error: {e:?}").into(), - ); - } - } + Some(flight) if flight.publish.qos == QoS::ExactlyOnce && !flight.released => { + state_mut.outbound.remove(&packet_id); + PubRecOutcome::Failed( + state_mut + .pending_pubcomps + .remove(&packet_id) + .map(|(callback, _)| callback), + ) } - _ => {} + _ => PubRecOutcome::Unknown, } + }; + match outcome { + PubRecOutcome::Release => send_ack(state, &Packet::PubRel(PubRelPacket::new(packet_id))), + PubRecOutcome::Failed(callback) => complete_flight(state, packet_id, reason_code, callback), + PubRecOutcome::Unknown => send_ack( + state, + &Packet::PubRel(PubRelPacket::new_with_reason( + packet_id, + ReasonCode::PacketIdentifierNotFound, + )), + ), } } -fn handle_pubcomp( - state: &Rc>, - packet_id: u16, - reason_code: mqtt5_protocol::protocol::v5::reason_codes::ReasonCode, -) { - let has_pending = state.borrow().pending_pubcomps.contains_key(&packet_id); - let actions = - mqtt5_protocol::qos2::handle_incoming_pubcomp(packet_id, reason_code, has_pending); - - for action in actions { - match action { - mqtt5_protocol::qos2::QoS2Action::CompleteFlow { packet_id } => { - if let Some((callback, _)) = state.borrow_mut().pending_pubcomps.remove(&packet_id) - { - let reason_code_js = JsValue::from_f64(f64::from(reason_code as u8)); - if let Err(e) = callback.call1(&JsValue::NULL, &reason_code_js) { - web_sys::console::error_1(&format!("PUBCOMP callback error: {e:?}").into()); - } - } - } - mqtt5_protocol::qos2::QoS2Action::ErrorFlow { - packet_id, - reason_code, - } => { - if let Some((callback, _)) = state.borrow_mut().pending_pubcomps.remove(&packet_id) - { - let reason_code_js = JsValue::from_f64(f64::from(reason_code as u8)); - if let Err(e) = callback.call1(&JsValue::NULL, &reason_code_js) { - web_sys::console::error_1( - &format!("QoS 2 error callback error: {e:?}").into(), - ); - } - } - } - _ => {} +fn handle_pubcomp(state: &Rc>, packet_id: u16, reason_code: ReasonCode) { + let callback = { + let mut state_mut = state.borrow_mut(); + let released = state_mut + .outbound + .get(&packet_id) + .is_some_and(|flight| flight.released); + if !released { + tracing::debug!(packet_id, "PUBCOMP for unknown packet identifier"); + return; } - } + state_mut.outbound.remove(&packet_id); + state_mut + .pending_pubcomps + .remove(&packet_id) + .map(|(callback, _)| callback) + }; + complete_flight(state, packet_id, reason_code, callback); } fn handle_pubrel(state: &Rc>, packet_id: u16) { - let has_pubrec = state.borrow().pending_pubrecs.contains_key(&packet_id); - let actions = mqtt5_protocol::qos2::handle_incoming_pubrel(packet_id, has_pubrec); - - for action in actions { - match action { - mqtt5_protocol::qos2::QoS2Action::RemoveIncomingPubRec { packet_id } => { - state.borrow_mut().pending_pubrecs.remove(&packet_id); - } - mqtt5_protocol::qos2::QoS2Action::SendPubComp { - packet_id, - reason_code, - } => { - let pubcomp = mqtt5_protocol::packet::pubcomp::PubCompPacket::new_with_reason( - packet_id, - reason_code, - ); - let mut buf = BytesMut::new(); - if let Err(e) = - encode_packet(&mqtt5_protocol::packet::Packet::PubComp(pubcomp), &mut buf) - { - web_sys::console::error_1(&format!("PUBCOMP encode error: {e}").into()); - continue; - } - - let writer_rc = state.borrow().writer.clone(); - if let Some(writer_rc) = writer_rc { - spawn_local(async move { - if let Err(e) = writer_rc.borrow_mut().write(&buf) { - web_sys::console::error_1(&format!("PUBCOMP send error: {e}").into()); - } - }); - } - } - _ => {} - } - } + let known = state.borrow_mut().awaiting_pubrel.remove(&packet_id); + let reason_code = if known { + ReasonCode::Success + } else { + ReasonCode::PacketIdentifierNotFound + }; + send_ack( + state, + &Packet::PubComp(PubCompPacket::new_with_reason(packet_id, reason_code)), + ); } diff --git a/crates/mqtt5-wasm/src/client/keepalive.rs b/crates/mqtt5-wasm/src/client/keepalive.rs index 404027e4..4b2b8adf 100644 --- a/crates/mqtt5-wasm/src/client/keepalive.rs +++ b/crates/mqtt5-wasm/src/client/keepalive.rs @@ -1,4 +1,3 @@ -use bytes::BytesMut; use mqtt5_protocol::packet::Packet; use mqtt5_protocol::time::Duration; use mqtt5_protocol::u128_to_u32_saturating; @@ -8,13 +7,19 @@ use std::rc::Rc; use wasm_bindgen_futures::spawn_local; use super::callbacks::handle_connection_lost; -use super::packet::encode_packet; +use super::packet::write_packet; use super::sleep_ms; use super::state::ClientState; pub fn spawn_keepalive_task(state: Rc>) { let keepalive_config = KeepaliveConfig::conservative(); - let generation = state.borrow().connection_generation; + let (generation, keep_alive) = { + let state_ref = state.borrow(); + (state_ref.connection_generation, state_ref.keep_alive) + }; + if keep_alive == 0 { + return; + } spawn_local(async move { loop { @@ -67,32 +72,12 @@ pub fn spawn_keepalive_task(state: Rc>) { break; } - let packet = Packet::PingReq; - let mut buf = BytesMut::new(); - if let Err(e) = encode_packet(&packet, &mut buf) { - web_sys::console::error_1(&format!("Ping encode error: {e}").into()); - continue; - } - state.borrow_mut().last_ping_sent = Some(js_sys::Date::now()); - let writer_rc = { - let state_ref = state.borrow(); - state_ref.writer.as_ref().map(Rc::clone) - }; - - match writer_rc { - Some(writer_rc) => match writer_rc.borrow_mut().write(&buf) { - Ok(()) => {} - Err(e) => { - let error_msg = format!("Ping send error: {e}"); - handle_connection_lost(&state, &error_msg); - break; - } - }, - None => { - break; - } + if let Err(e) = write_packet(&state, &Packet::PingReq) { + let error_msg = format!("Ping send error: {e}"); + handle_connection_lost(&state, &error_msg); + break; } } }); diff --git a/crates/mqtt5-wasm/src/client/mod.rs b/crates/mqtt5-wasm/src/client/mod.rs index 8acf64df..8d2de9d7 100644 --- a/crates/mqtt5-wasm/src/client/mod.rs +++ b/crates/mqtt5-wasm/src/client/mod.rs @@ -1,7 +1,9 @@ mod callbacks; +mod connection; mod connectivity; mod handlers; mod keepalive; +mod outbound; mod packet; mod qos; mod reader; @@ -11,35 +13,33 @@ mod state; use crate::config::{ WasmConnectOptions, WasmPublishOptions, WasmReconnectOptions, WasmSubscribeOptions, }; -use crate::decoder::read_packet; -use crate::transport::{WasmReader, WasmTransportType}; -use bytes::BytesMut; +use crate::transport::WasmTransportType; use mqtt5_protocol::packet::connect::ConnectPacket; +use mqtt5_protocol::packet::disconnect::DisconnectPacket; use mqtt5_protocol::packet::publish::PublishPacket; -use mqtt5_protocol::packet::subscribe::SubscribePacket; +use mqtt5_protocol::packet::subscribe::{SubscribePacket, TopicFilter}; use mqtt5_protocol::packet::unsubscribe::UnsubscribePacket; use mqtt5_protocol::packet::Packet; -use mqtt5_protocol::protocol::v5::properties::Properties; +use mqtt5_protocol::protocol::v5::properties::{Properties, PropertyId, PropertyValue}; use mqtt5_protocol::strip_shared_subscription_prefix; use mqtt5_protocol::QoS; -use mqtt5_protocol::Transport; use std::cell::RefCell; use std::rc::Rc; use wasm_bindgen::prelude::*; use wasm_bindgen_futures::JsFuture; use web_sys::MessagePort; -use callbacks::{drain_pending_callbacks, trigger_disconnect_callback}; -use keepalive::spawn_keepalive_task; -use packet::encode_packet; -use qos::{await_ack_promises, create_ack_promises, spawn_qos2_cleanup_task}; -use reader::spawn_packet_reader; +use callbacks::{close_network_connection, end_connection_state, trigger_disconnect_callback}; +use connection::establish; +use outbound::{check_publish, check_subscribe, check_unsubscribe}; +use packet::write_packet; +use qos::{abandon_flight, await_ack_promises, create_ack_promises, reserve_flight}; use state::{ClientState, StoredConnectOptions}; #[wasm_bindgen] extern "C" { #[wasm_bindgen(js_name = "setTimeout")] - fn set_timeout(handler: &js_sys::Function, timeout: i32) -> i32; + fn set_timeout(handler: &js_sys::Function, timeout: i32) -> JsValue; } pub async fn sleep_ms(millis: u32) { @@ -59,6 +59,15 @@ pub struct RustMessage { type RustCallback = Rc; +enum AckSink { + Promise, + Callback(js_sys::Function), +} + +fn js_error(message: impl AsRef) -> JsValue { + JsValue::from_str(message.as_ref()) +} + #[wasm_bindgen(js_name = "MqttClient")] pub struct WasmMqttClient { state: Rc>, @@ -67,11 +76,11 @@ pub struct WasmMqttClient { #[wasm_bindgen(js_class = "MqttClient")] impl WasmMqttClient { #[wasm_bindgen(constructor)] - #[allow(clippy::must_use_candidate, non_snake_case)] - pub fn new(clientId: String) -> Self { + #[must_use] + pub fn new(#[wasm_bindgen(js_name = clientId)] client_id: String) -> Self { console_error_panic_hook::set_once(); - let state = Rc::new(RefCell::new(ClientState::new(clientId))); + let state = Rc::new(RefCell::new(ClientState::new(client_id))); let (online_fn, offline_fn) = connectivity::register_connectivity_listeners(&state); { let mut s = state.borrow_mut(); @@ -99,10 +108,10 @@ impl WasmMqttClient { ) -> Result<(), JsValue> { let lower = url.to_ascii_lowercase(); if !lower.starts_with("ws://") && !lower.starts_with("wss://") { - return Err(JsValue::from_str("URL must start with ws:// or wss://")); + return Err(js_error("URL must start with ws:// or wss://")); } if lower.starts_with("ws://") && (config.username.is_some() || config.password.is_some()) { - web_sys::console::warn_1(&"Credentials sent over unencrypted ws:// connection".into()); + tracing::warn!("credentials sent over an unencrypted ws:// connection"); } self.state.borrow_mut().last_url = Some(url.to_string()); @@ -139,11 +148,13 @@ impl WasmMqttClient { /// # Errors /// Returns an error if connection fails. #[wasm_bindgen(js_name = "connectBroadcastChannel")] - #[allow(non_snake_case)] - pub async fn connect_broadcast_channel(&self, channelName: &str) -> Result<(), JsValue> { + pub async fn connect_broadcast_channel( + &self, + #[wasm_bindgen(js_name = channelName)] channel_name: &str, + ) -> Result<(), JsValue> { let config = WasmConnectOptions::default(); let transport = WasmTransportType::BroadcastChannel( - crate::transport::broadcast::BroadcastChannelTransport::new(channelName), + crate::transport::broadcast::BroadcastChannelTransport::new(channel_name), ); self.connect_with_transport_and_config(transport, &config) .await @@ -151,177 +162,34 @@ impl WasmMqttClient { async fn connect_with_transport_and_config( &self, - mut transport: WasmTransportType, + transport: WasmTransportType, config: &WasmConnectOptions, ) -> Result<(), JsValue> { - { - let state_ref = self.state.borrow(); - if state_ref.connected { - return Err(JsValue::from_str("Already connected")); + let stored = StoredConnectOptions::from(config); + let client_id = { + let mut state = self.state.borrow_mut(); + if state.connected { + return Err(js_error("Already connected")); } - if state_ref.reconnecting { - return Err(JsValue::from_str("Reconnection in progress")); + if state.reconnecting { + return Err(js_error("Reconnection in progress")); } - } - - transport - .connect() - .await - .map_err(|e| JsValue::from_str(&format!("Transport connection failed: {e}")))?; - - let client_id = self.state.borrow().client_id.clone(); - - { - let mut state = self.state.borrow_mut(); - state.keep_alive = config.keep_alive; - state.protocol_version = config.protocol_version; - state.last_options = Some(StoredConnectOptions::from(config)); + state.last_options = Some(stored.clone()); state.user_initiated_disconnect = false; state.reconnect_attempt = 0; - #[cfg(feature = "codec")] - { - state.codec_registry.clone_from(&config.codec_registry); - } - } - - let packet = Packet::Connect(Box::new(build_connect_packet(client_id, config))); - let mut buf = BytesMut::new(); - encode_packet(&packet, &mut buf) - .map_err(|e| JsValue::from_str(&format!("Packet encoding failed: {e}")))?; + state.client_id.clone() + }; - transport - .write(&buf) + let connect = build_connect_packet(client_id, config); + establish(&self.state, transport, connect, &stored) .await - .map_err(|e| JsValue::from_str(&format!("Write failed: {e}")))?; - - if let Some(method) = &config.authentication_method { - self.state.borrow_mut().auth_method = Some(method.clone()); - } - - let (reader, writer) = transport - .into_split() - .map_err(|e| JsValue::from_str(&format!("Transport split failed: {e}")))?; - - let writer_rc = Rc::new(RefCell::new(writer)); - self.state.borrow_mut().writer = Some(Rc::clone(&writer_rc)); - - self.handle_connect_response(reader).await - } - - async fn handle_connect_response(&self, mut reader: WasmReader) -> Result<(), JsValue> { - loop { - let packet = read_packet(&mut reader) - .await - .map_err(|e| JsValue::from_str(&format!("Packet read failed: {e}")))?; - - match packet { - Packet::ConnAck(connack) => { - let reason_code = connack.reason_code as u8; - let session_present = connack.session_present; - - if reason_code != 0 { - if reason_code == 0x9C || reason_code == 0x9D { - if let Some(server_ref) = connack.properties.get_server_reference() { - let obj = js_sys::Object::new(); - js_sys::Reflect::set( - &obj, - &JsValue::from_str("type"), - &JsValue::from_str("redirect"), - ) - .ok(); - js_sys::Reflect::set( - &obj, - &JsValue::from_str("url"), - &JsValue::from_str(server_ref), - ) - .ok(); - return Err(obj.into()); - } - } - return Err(JsValue::from_str(&format!( - "Connection rejected: {}", - connack_error_description(reason_code) - ))); - } - - { - let mut state_mut = self.state.borrow_mut(); - state_mut.connected = true; - state_mut.connection_generation = - state_mut.connection_generation.wrapping_add(1); - } - - spawn_packet_reader(Rc::clone(&self.state), reader); - spawn_keepalive_task(Rc::clone(&self.state)); - spawn_qos2_cleanup_task(Rc::clone(&self.state)); - - let callback = self.state.borrow().on_connect.clone(); - if let Some(callback) = callback { - let reason_code_js = JsValue::from_f64(f64::from(reason_code)); - let session_present_js = JsValue::from_bool(session_present); - - if let Err(e) = - callback.call2(&JsValue::NULL, &reason_code_js, &session_present_js) - { - web_sys::console::error_1( - &format!("onConnect callback error: {e:?}").into(), - ); - } - } - - return Ok(()); - } - Packet::Auth(auth) => { - let auth_reason = auth.reason_code; - if auth_reason - == mqtt5_protocol::protocol::v5::reason_codes::ReasonCode::ContinueAuthentication - { - let callback = self.state.borrow().on_auth_challenge.clone(); - if let Some(callback) = callback { - let auth_method = auth - .properties - .get_authentication_method() - .cloned() - .unwrap_or_default(); - let auth_data = auth.properties.get_authentication_data(); - - let method_js = JsValue::from_str(&auth_method); - let data_js = if let Some(data) = auth_data { - js_sys::Uint8Array::from(data).into() - } else { - JsValue::NULL - }; - - if let Err(e) = callback.call2(&JsValue::NULL, &method_js, &data_js) { - web_sys::console::error_1( - &format!("onAuthChallenge callback error: {e:?}").into(), - ); - } - } else { - return Err(JsValue::from_str( - "AUTH challenge received but no on_auth_challenge callback set", - )); - } - } else { - return Err(JsValue::from_str(&format!( - "Unexpected AUTH reason code: {auth_reason:?}" - ))); - } - } - _ => { - return Err(JsValue::from_str(&format!( - "Expected CONNACK or AUTH, received: {packet:?}" - ))); - } - } - } + .map(|_| ()) + .map_err(JsValue::from) } /// # Errors /// Returns an error if not connected or publish fails. pub async fn publish(&self, topic: &str, payload: &[u8]) -> Result<(), JsValue> { - self.ensure_connected().await?; - let protocol_version = self.state.borrow().protocol_version; let publish_packet = PublishPacket { dup: false, @@ -334,8 +202,9 @@ impl WasmMqttClient { protocol_version, stream_id: None, }; - - self.send_packet(&Packet::Publish(publish_packet)) + self.dispatch_publish(publish_packet, AckSink::Promise) + .await + .map(|_| ()) } /// # Errors @@ -347,16 +216,7 @@ impl WasmMqttClient { payload: &[u8], options: &WasmPublishOptions, ) -> Result<(), JsValue> { - self.ensure_connected().await?; - let qos = options.to_qos(); - let packet_id = if qos == QoS::AtMostOnce { - None - } else { - Some(self.state.borrow_mut().packet_id.next()) - }; - - let (puback_promise, pubcomp_promise) = create_ack_promises(&self.state, qos, packet_id); #[cfg(feature = "codec")] let (final_payload, codec_content_type) = { @@ -384,13 +244,10 @@ impl WasmMqttClient { if let Some(ct) = codec_content_type { if properties - .add( - mqtt5_protocol::protocol::v5::properties::PropertyId::ContentType, - mqtt5_protocol::protocol::v5::properties::PropertyValue::Utf8String(ct), - ) + .add(PropertyId::ContentType, PropertyValue::Utf8String(ct)) .is_err() { - web_sys::console::warn_1(&"Failed to add codec content type property".into()); + tracing::warn!("failed to add codec content type property"); } } @@ -399,14 +256,17 @@ impl WasmMqttClient { qos, retain: options.retain, topic_name: topic.to_string(), - packet_id, + packet_id: None, properties, payload: final_payload.into(), protocol_version, stream_id: None, }; - self.send_packet(&Packet::Publish(publish_packet))?; + let packet_id = self + .dispatch_publish(publish_packet, AckSink::Promise) + .await?; + let (puback_promise, pubcomp_promise) = create_ack_promises(&self.state, qos, packet_id); await_ack_promises(puback_promise, pubcomp_promise).await } @@ -419,29 +279,8 @@ impl WasmMqttClient { payload: &[u8], callback: js_sys::Function, ) -> Result { - self.ensure_connected().await?; - - let packet_id = self.state.borrow_mut().packet_id.next(); - self.state - .borrow_mut() - .pending_pubacks - .insert(packet_id, callback); - - let protocol_version = self.state.borrow().protocol_version; - let publish_packet = PublishPacket { - dup: false, - qos: QoS::AtLeastOnce, - retain: false, - topic_name: topic.to_string(), - packet_id: Some(packet_id), - properties: Properties::default(), - payload: payload.to_vec().into(), - protocol_version, - stream_id: None, - }; - - self.send_packet(&Packet::Publish(publish_packet))?; - Ok(packet_id) + self.publish_with_callback(topic, payload, QoS::AtLeastOnce, callback) + .await } /// # Errors @@ -453,60 +292,20 @@ impl WasmMqttClient { payload: &[u8], callback: js_sys::Function, ) -> Result { - self.ensure_connected().await?; - - let packet_id = self.state.borrow_mut().packet_id.next(); - let now = js_sys::Date::now(); - self.state - .borrow_mut() - .pending_pubcomps - .insert(packet_id, (callback, now)); - - let protocol_version = self.state.borrow().protocol_version; - let publish_packet = PublishPacket { - dup: false, - qos: QoS::ExactlyOnce, - retain: false, - topic_name: topic.to_string(), - packet_id: Some(packet_id), - properties: Properties::default(), - payload: payload.to_vec().into(), - protocol_version, - stream_id: None, - }; - - self.send_packet(&Packet::Publish(publish_packet))?; - Ok(packet_id) + self.publish_with_callback(topic, payload, QoS::ExactlyOnce, callback) + .await } /// # Errors /// Returns an error if not connected or subscribe fails. - #[allow(clippy::unused_async, clippy::unused_async_trait_impl)] pub async fn subscribe(&self, topic: &str) -> Result { - if !self.state.borrow().connected { - return Err(JsValue::from_str("Not connected")); - } - - let packet_id = self.state.borrow_mut().packet_id.next(); - - let protocol_version = self.state.borrow().protocol_version; - let subscribe_packet = SubscribePacket { - packet_id, - properties: Properties::default(), - filters: vec![mqtt5_protocol::packet::subscribe::TopicFilter::new( - topic, - QoS::AtMostOnce, - )], - protocol_version, - }; - - self.send_packet(&Packet::Subscribe(subscribe_packet))?; - Ok(packet_id) + let filter = TopicFilter::new(topic, QoS::AtMostOnce); + self.send_subscribe(filter, Properties::default(), None) + .await } /// # Errors /// Returns an error if not connected or subscribe fails. - #[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)] #[wasm_bindgen(js_name = "subscribeWithOptions")] pub async fn subscribe_with_options( &self, @@ -514,19 +313,7 @@ impl WasmMqttClient { callback: js_sys::Function, options: &WasmSubscribeOptions, ) -> Result { - if !self.state.borrow().connected { - return Err(JsValue::from_str("Not connected")); - } - - let packet_id = self.state.borrow_mut().packet_id.next(); - let actual_filter = strip_shared_subscription_prefix(topic); - self.state - .borrow_mut() - .subscriptions - .insert(actual_filter.to_string(), callback); - - let mut topic_filter = - mqtt5_protocol::packet::subscribe::TopicFilter::new(topic, options.to_qos()); + let mut topic_filter = TopicFilter::new(topic, options.to_qos()); topic_filter.options.no_local = options.no_local; topic_filter.options.retain_as_published = options.retain_as_published; topic_filter.options.retain_handling = match options.retain_handling { @@ -539,52 +326,35 @@ impl WasmMqttClient { if let Some(id) = options.subscription_identifier { if properties .add( - mqtt5_protocol::protocol::v5::properties::PropertyId::SubscriptionIdentifier, - mqtt5_protocol::protocol::v5::properties::PropertyValue::VariableByteInteger( - id, - ), + PropertyId::SubscriptionIdentifier, + PropertyValue::VariableByteInteger(id), ) .is_err() { - web_sys::console::warn_1(&"Failed to add subscription identifier property".into()); + tracing::warn!("failed to add subscription identifier property"); } } - let protocol_version = self.state.borrow().protocol_version; - let properties = if protocol_version == 5 { - properties - } else { - Properties::default() - }; - let subscribe_packet = SubscribePacket { - packet_id, - properties, - filters: vec![topic_filter], - protocol_version, - }; - - self.send_packet(&Packet::Subscribe(subscribe_packet))?; + let packet_id = self + .send_subscribe(topic_filter, properties, Some(callback)) + .await?; let state = Rc::clone(&self.state); let promise = js_sys::Promise::new(&mut move |resolve, _reject| { - state - .borrow_mut() - .pending_subacks - .insert(packet_id, resolve); + if let Some(slot) = state.borrow_mut().pending_subacks.get_mut(&packet_id) { + *slot = Some(resolve); + } }); let result = JsFuture::from(promise).await?; let reason_codes = js_sys::Array::from(&result); - - if reason_codes.length() > 0 { - let first_code = reason_codes.get(0).as_f64().unwrap_or(0.0) as u8; - if first_code >= 0x80 { - let actual_filter = strip_shared_subscription_prefix(topic); - self.state.borrow_mut().subscriptions.remove(actual_filter); - return Err(JsValue::from_str(&format!( - "Subscribe rejected with reason code: {first_code}" - ))); - } + let first_code = reason_codes.get(0).as_f64().unwrap_or(0.0); + if first_code >= 128.0 { + let actual_filter = strip_shared_subscription_prefix(topic); + self.state.borrow_mut().subscriptions.remove(actual_filter); + return Err(js_error(format!( + "Subscribe rejected with reason code: {first_code}" + ))); } Ok(packet_id) @@ -592,99 +362,63 @@ impl WasmMqttClient { /// # Errors /// Returns an error if not connected or subscribe fails. - #[allow(clippy::unused_async, clippy::unused_async_trait_impl)] #[wasm_bindgen(js_name = "subscribeWithCallback")] pub async fn subscribe_with_callback( &self, topic: &str, callback: js_sys::Function, ) -> Result { - if !self.state.borrow().connected { - return Err(JsValue::from_str("Not connected")); - } - - let packet_id = self.state.borrow_mut().packet_id.next(); - let actual_filter = strip_shared_subscription_prefix(topic); - self.state - .borrow_mut() - .subscriptions - .insert(actual_filter.to_string(), callback); - - let protocol_version = self.state.borrow().protocol_version; - let subscribe_packet = SubscribePacket { - packet_id, - properties: Properties::default(), - filters: vec![mqtt5_protocol::packet::subscribe::TopicFilter::new( - topic, - QoS::AtMostOnce, - )], - protocol_version, - }; - - self.send_packet(&Packet::Subscribe(subscribe_packet))?; - Ok(packet_id) + let filter = TopicFilter::new(topic, QoS::AtMostOnce); + self.send_subscribe(filter, Properties::default(), Some(callback)) + .await } /// # Errors /// Returns an error if not connected or unsubscribe fails. - #[allow(clippy::unused_async, clippy::unused_async_trait_impl)] pub async fn unsubscribe(&self, topic: &str) -> Result { - if !self.state.borrow().connected { - return Err(JsValue::from_str("Not connected")); - } - - let packet_id = self.state.borrow_mut().packet_id.next(); - self.state.borrow_mut().subscriptions.remove(topic); + self.ensure_connected().await?; let protocol_version = self.state.borrow().protocol_version; - let unsubscribe_packet = UnsubscribePacket { - packet_id, + let mut unsubscribe_packet = UnsubscribePacket { + packet_id: 0, properties: Properties::default(), filters: vec![topic.to_string()], protocol_version, }; + check_unsubscribe(&self.state.borrow(), &unsubscribe_packet).map_err(js_error)?; + + let packet_id = self.reserve_packet_id()?; + unsubscribe_packet.packet_id = packet_id; + self.state.borrow_mut().pending_unsubacks.insert(packet_id); - self.send_packet(&Packet::Unsubscribe(unsubscribe_packet))?; + if let Err(e) = write_packet(&self.state, &Packet::Unsubscribe(unsubscribe_packet)) { + self.state.borrow_mut().pending_unsubacks.remove(&packet_id); + return Err(js_error(e)); + } + self.state + .borrow_mut() + .subscriptions + .remove(strip_shared_subscription_prefix(topic)); Ok(packet_id) } /// # Errors /// Returns an error if disconnect fails. pub async fn disconnect(&self) -> Result<(), JsValue> { - let disconnect_packet = mqtt5_protocol::packet::disconnect::DisconnectPacket { - reason_code: mqtt5_protocol::protocol::v5::reason_codes::ReasonCode::Success, - properties: Properties::default(), + self.ensure_not_borrowed().await; + let connected = { + let mut state = self.state.borrow_mut(); + state.user_initiated_disconnect = true; + state.connected }; - let packet = Packet::Disconnect(disconnect_packet); - let mut buf = BytesMut::new(); - encode_packet(&packet, &mut buf) - .map_err(|e| JsValue::from_str(&format!("DISCONNECT packet encoding failed: {e}")))?; - - let writer_rc = loop { - match self.state.try_borrow_mut() { - Ok(mut state) => { - state.connected = false; - state.user_initiated_disconnect = true; - state.connection_generation = state.connection_generation.wrapping_add(1); - break state.writer.take(); - } - Err(_) => { - sleep_ms(10).await; - } + if connected { + let disconnect = DisconnectPacket::normal(); + if let Err(e) = write_packet(&self.state, &Packet::Disconnect(disconnect)) { + tracing::warn!(error = %e, "DISCONNECT not sent"); } - }; - - if let Some(writer_rc) = writer_rc { - let mut writer = writer_rc.borrow_mut(); - writer - .write(&buf) - .map_err(|e| JsValue::from_str(&format!("DISCONNECT packet send failed: {e}")))?; - writer - .close() - .map_err(|e| JsValue::from_str(&format!("Close failed: {e}")))?; } - - drain_pending_callbacks(&mut self.state.borrow_mut()); + close_network_connection(&self.state); + end_connection_state(&self.state); trigger_disconnect_callback(&self.state); Ok(()) } @@ -764,14 +498,16 @@ impl WasmMqttClient { /// # Errors /// Returns an error if no auth method is set or send fails. #[wasm_bindgen(js_name = "respondAuth")] - #[allow(non_snake_case)] - pub fn respond_auth(&self, authData: &[u8]) -> Result<(), JsValue> { + pub fn respond_auth( + &self, + #[wasm_bindgen(js_name = authData)] auth_data: &[u8], + ) -> Result<(), JsValue> { let auth_method = self .state .borrow() .auth_method .clone() - .ok_or_else(|| JsValue::from_str("No auth method set"))?; + .ok_or_else(|| js_error("No auth method set"))?; let mut auth_packet = mqtt5_protocol::packet::auth::AuthPacket::new( mqtt5_protocol::protocol::v5::reason_codes::ReasonCode::ContinueAuthentication, @@ -781,44 +517,160 @@ impl WasmMqttClient { .set_authentication_method(auth_method); auth_packet .properties - .set_authentication_data(authData.to_vec().into()); + .set_authentication_data(auth_data.to_vec().into()); - self.send_packet(&Packet::Auth(auth_packet)) + write_packet(&self.state, &Packet::Auth(auth_packet)).map_err(js_error) + } + + async fn ensure_not_borrowed(&self) { + while self.state.try_borrow_mut().is_err() { + sleep_ms(10).await; + } } async fn ensure_connected(&self) -> Result<(), JsValue> { - loop { - match self.state.try_borrow() { - Ok(state) => { - if !state.connected { - return Err(JsValue::from_str("Not connected")); - } - return Ok(()); - } - Err(_) => { - sleep_ms(10).await; + self.ensure_not_borrowed().await; + if self.state.borrow().connected { + Ok(()) + } else { + Err(js_error("Not connected")) + } + } + + fn reserve_packet_id(&self) -> Result { + self.state + .borrow() + .allocate_packet_id() + .ok_or_else(|| js_error("No packet identifier available")) + } + + async fn publish_with_callback( + &self, + topic: &str, + payload: &[u8], + qos: QoS, + callback: js_sys::Function, + ) -> Result { + let protocol_version = self.state.borrow().protocol_version; + let publish_packet = PublishPacket { + dup: false, + qos, + retain: false, + topic_name: topic.to_string(), + packet_id: None, + properties: Properties::default(), + payload: payload.to_vec().into(), + protocol_version, + stream_id: None, + }; + self.dispatch_publish(publish_packet, AckSink::Callback(callback)) + .await? + .ok_or_else(|| js_error("QoS 0 publish has no packet identifier")) + } + + async fn dispatch_publish( + &self, + mut publish: PublishPacket, + sink: AckSink, + ) -> Result, JsValue> { + self.ensure_connected().await?; + check_publish(&self.state.borrow(), &publish).map_err(js_error)?; + + let packet_id = if publish.qos == QoS::AtMostOnce { + None + } else { + let packet_id = reserve_flight(&self.state, &publish).await?; + if let Err(e) = check_publish(&self.state.borrow(), &publish) { + abandon_flight(&self.state, packet_id); + return Err(js_error(e)); + } + Some(packet_id) + }; + publish.packet_id = packet_id; + + if let (Some(packet_id), AckSink::Callback(callback)) = (packet_id, sink) { + let mut state = self.state.borrow_mut(); + if publish.qos == QoS::ExactlyOnce { + state + .pending_pubcomps + .insert(packet_id, (callback, js_sys::Date::now())); + } else { + state.pending_pubacks.insert(packet_id, callback); + } + } + + let alias_mapping = publish + .topic_alias() + .filter(|_| !publish.topic_name.is_empty()) + .map(|alias| (alias, publish.topic_name.clone())); + + if let Err(e) = write_packet(&self.state, &Packet::Publish(publish)) { + if let Some(packet_id) = packet_id { + { + let mut state = self.state.borrow_mut(); + state.pending_pubacks.remove(&packet_id); + state.pending_pubcomps.remove(&packet_id); } + abandon_flight(&self.state, packet_id); } + return Err(js_error(e)); } + + if let Some((alias, topic)) = alias_mapping { + if let Err(e) = self + .state + .borrow_mut() + .outbound_aliases + .register_alias(alias, &topic) + { + tracing::warn!(alias, topic, error = %e, "outbound Topic Alias not recorded"); + } + } + + Ok(packet_id) } - fn send_packet(&self, packet: &Packet) -> Result<(), JsValue> { - let mut buf = BytesMut::new(); - encode_packet(packet, &mut buf) - .map_err(|e| JsValue::from_str(&format!("Packet encoding failed: {e}")))?; + async fn send_subscribe( + &self, + filter: TopicFilter, + properties: Properties, + callback: Option, + ) -> Result { + self.ensure_connected().await?; - let writer_rc = self - .state - .borrow() - .writer - .clone() - .ok_or_else(|| JsValue::from_str("Writer disconnected"))?; + let protocol_version = self.state.borrow().protocol_version; + let properties = if protocol_version == 5 { + properties + } else { + Properties::default() + }; + let topic = filter.filter.clone(); + let mut subscribe_packet = SubscribePacket { + packet_id: 0, + properties, + filters: vec![filter], + protocol_version, + }; + check_subscribe(&self.state.borrow(), &subscribe_packet).map_err(js_error)?; - let result = writer_rc - .borrow_mut() - .write(&buf) - .map_err(|e| JsValue::from_str(&format!("Write failed: {e}"))); - result + let packet_id = self.reserve_packet_id()?; + subscribe_packet.packet_id = packet_id; + { + let mut state = self.state.borrow_mut(); + state.pending_subacks.insert(packet_id, None); + if let Some(callback) = callback { + state.subscriptions.insert( + strip_shared_subscription_prefix(&topic).to_string(), + callback, + ); + } + } + + if let Err(e) = write_packet(&self.state, &Packet::Subscribe(subscribe_packet)) { + self.state.borrow_mut().pending_subacks.remove(&packet_id); + return Err(js_error(e)); + } + Ok(packet_id) } } @@ -850,29 +702,6 @@ fn build_connect_packet(client_id: String, config: &WasmConnectOptions) -> Conne } } -fn connack_error_description(reason_code: u8) -> &'static str { - match reason_code { - 0x80 => "Unspecified error", - 0x81 => "Malformed packet", - 0x82 => "Protocol error", - 0x83 => "Implementation specific error", - 0x84 => "Unsupported protocol version", - 0x85 => "Client identifier not valid", - 0x86 => "Bad username or password", - 0x87 => "Not authorized", - 0x88 => "Server unavailable", - 0x89 => "Server busy", - 0x8A => "Banned", - 0x8C => "Bad authentication method", - 0x90 => "Topic name invalid", - 0x97 => "Quota exceeded", - 0x9C => "Use another server", - 0x9D => "Server moved", - 0x9F => "Connection rate exceeded", - _ => "Unknown error", - } -} - impl WasmMqttClient { /// # Errors /// Returns an error if not connected or subscribe fails. @@ -888,7 +717,6 @@ impl WasmMqttClient { /// # Errors /// Returns an error if not connected or subscribe fails. - #[allow(clippy::unused_async, clippy::unused_async_trait_impl)] pub async fn subscribe_with_callback_internal_opts( &self, topic: &str, @@ -896,31 +724,19 @@ impl WasmMqttClient { no_local: bool, callback: Box, ) -> Result { - if !self.state.borrow().connected { - return Err(JsValue::from_str("Not connected")); - } - - let packet_id = self.state.borrow_mut().packet_id.next(); - let actual_filter = strip_shared_subscription_prefix(topic); - self.state - .borrow_mut() - .rust_subscriptions - .insert(actual_filter.to_string(), Rc::new(callback)); - let mut options = mqtt5_protocol::packet::subscribe::SubscriptionOptions::new(qos); options.no_local = no_local; - - let protocol_version = self.state.borrow().protocol_version; - let subscribe_packet = SubscribePacket { - packet_id, - properties: Properties::default(), - filters: vec![ - mqtt5_protocol::packet::subscribe::TopicFilter::with_options(topic, options), - ], - protocol_version, - }; - - self.send_packet(&Packet::Subscribe(subscribe_packet))?; + let packet_id = self + .send_subscribe( + TopicFilter::with_options(topic, options), + Properties::default(), + None, + ) + .await?; + self.state.borrow_mut().rust_subscriptions.insert( + strip_shared_subscription_prefix(topic).to_string(), + Rc::new(callback), + ); Ok(packet_id) } @@ -932,22 +748,11 @@ impl WasmMqttClient { payload: &[u8], qos: QoS, ) -> Result<(), JsValue> { - if !self.state.borrow().connected { - return Err(JsValue::from_str("Not connected")); - } - - let packet_id = if qos == QoS::AtMostOnce { - None - } else { - Some(self.state.borrow_mut().packet_id.next()) - }; - + let publish_packet = PublishPacket::new(topic.to_string(), payload.to_vec(), qos); + let packet_id = self + .dispatch_publish(publish_packet, AckSink::Promise) + .await?; let (puback_promise, pubcomp_promise) = create_ack_promises(&self.state, qos, packet_id); - - let mut publish_packet = PublishPacket::new(topic.to_string(), payload.to_vec(), qos); - publish_packet.packet_id = packet_id; - - self.send_packet(&Packet::Publish(publish_packet))?; await_ack_promises(puback_promise, pubcomp_promise).await } } diff --git a/crates/mqtt5-wasm/src/client/outbound.rs b/crates/mqtt5-wasm/src/client/outbound.rs new file mode 100644 index 00000000..9f47d694 --- /dev/null +++ b/crates/mqtt5-wasm/src/client/outbound.rs @@ -0,0 +1,120 @@ +use mqtt5_protocol::constants::variable_byte::MAX_VALUE as MAX_SUBSCRIPTION_IDENTIFIER; +use mqtt5_protocol::packet::publish::PublishPacket; +use mqtt5_protocol::packet::subscribe::SubscribePacket; +use mqtt5_protocol::packet::unsubscribe::UnsubscribePacket; +use mqtt5_protocol::packet::Packet; +use mqtt5_protocol::protocol::v5::properties::PropertyId; +use mqtt5_protocol::validation::{ + parse_shared_subscription, validate_subscription_filter, validate_topic_name, +}; +use mqtt5_protocol::QoS; + +use super::packet::encode_checked; +use super::state::ClientState; + +pub fn check_publish(state: &ClientState, publish: &PublishPacket) -> Result<(), String> { + let alias = publish.topic_alias(); + match (publish.topic_name.is_empty(), alias) { + (true, None) => return Err("A zero-length Topic Name requires a Topic Alias".to_string()), + (true, Some(_)) => {} + (false, _) => validate_topic_name(&publish.topic_name).map_err(|e| e.to_string())?, + } + if let Some(alias) = alias { + check_topic_alias(state, &publish.topic_name, alias)?; + } + if let Some(response_topic) = response_topic(publish) { + validate_topic_name(response_topic).map_err(|e| format!("Invalid Response Topic: {e}"))?; + } + if publish + .properties + .contains(PropertyId::SubscriptionIdentifier) + { + return Err("A client PUBLISH must not contain a Subscription Identifier".to_string()); + } + if publish.qos as u8 > state.server.maximum_qos as u8 { + return Err(format!( + "QoS {} exceeds the server Maximum QoS {}", + publish.qos as u8, state.server.maximum_qos as u8 + )); + } + if publish.retain && !state.server.retain_available { + return Err("The server does not support retained messages".to_string()); + } + let mut sized = publish.clone(); + if sized.qos != QoS::AtMostOnce { + sized.packet_id = Some(sized.packet_id.unwrap_or(u16::MAX)); + } + encode_checked(&Packet::Publish(sized), state.server.maximum_packet_size)?; + Ok(()) +} + +fn response_topic(publish: &PublishPacket) -> Option<&str> { + match publish.properties.get(PropertyId::ResponseTopic) { + Some(mqtt5_protocol::protocol::v5::properties::PropertyValue::Utf8String(topic)) => { + Some(topic.as_str()) + } + _ => None, + } +} + +fn check_topic_alias(state: &ClientState, topic: &str, alias: u16) -> Result<(), String> { + if alias == 0 { + return Err("Topic Alias 0 is not permitted".to_string()); + } + if alias > state.server.topic_alias_maximum { + return Err(format!( + "Topic Alias {alias} exceeds the server Topic Alias Maximum {}", + state.server.topic_alias_maximum + )); + } + if topic.is_empty() && state.outbound_aliases.get_topic(alias).is_none() { + return Err(format!( + "Topic Alias {alias} has no mapping on this connection" + )); + } + Ok(()) +} + +pub fn check_subscribe(state: &ClientState, packet: &SubscribePacket) -> Result<(), String> { + let identifiers = packet.properties.subscription_identifiers(); + if let Some(invalid) = identifiers + .iter() + .find(|id| !(1..=MAX_SUBSCRIPTION_IDENTIFIER).contains(*id)) + { + return Err(format!( + "Subscription Identifier {invalid} is outside 1..={MAX_SUBSCRIPTION_IDENTIFIER}" + )); + } + if !identifiers.is_empty() && !state.server.subscriptions.identifiers { + return Err("The server does not support Subscription Identifiers".to_string()); + } + for filter in &packet.filters { + validate_subscription_filter(&filter.filter).map_err(|e| e.to_string())?; + let (topic_filter, share_name) = parse_shared_subscription(&filter.filter); + if share_name.is_some() && !state.server.subscriptions.shared { + return Err("The server does not support Shared Subscriptions".to_string()); + } + if share_name.is_some() && filter.options.no_local { + return Err("No Local must not be set on a Shared Subscription".to_string()); + } + if topic_filter.contains(['+', '#']) && !state.server.subscriptions.wildcards { + return Err("The server does not support Wildcard Subscriptions".to_string()); + } + } + encode_checked( + &Packet::Subscribe(packet.clone()), + state.server.maximum_packet_size, + )?; + Ok(()) +} + +pub fn check_unsubscribe(state: &ClientState, packet: &UnsubscribePacket) -> Result<(), String> { + for filter in &packet.filters { + validate_subscription_filter(filter).map_err(|e| e.to_string())?; + } + encode_checked( + &Packet::Unsubscribe(packet.clone()), + state.server.maximum_packet_size, + )?; + Ok(()) +} diff --git a/crates/mqtt5-wasm/src/client/packet.rs b/crates/mqtt5-wasm/src/client/packet.rs index 8b7c258b..23001855 100644 --- a/crates/mqtt5-wasm/src/client/packet.rs +++ b/crates/mqtt5-wasm/src/client/packet.rs @@ -1,11 +1,16 @@ use bytes::BytesMut; use mqtt5_protocol::error::Result; use mqtt5_protocol::packet::{MqttPacket, Packet}; +use std::cell::RefCell; +use std::rc::Rc; + +use super::state::ClientState; pub fn encode_packet(packet: &Packet, buf: &mut BytesMut) -> Result<()> { match packet { Packet::Connect(p) => p.encode(buf), Packet::Publish(p) => p.encode(buf), + Packet::PubAck(p) => p.encode(buf), Packet::PubRec(p) => p.encode(buf), Packet::PubRel(p) => p.encode(buf), Packet::PubComp(p) => p.encode(buf), @@ -19,3 +24,39 @@ pub fn encode_packet(packet: &Packet, buf: &mut BytesMut) -> Result<()> { ))), } } + +pub fn encode_checked( + packet: &Packet, + maximum_packet_size: Option, +) -> std::result::Result { + let mut buf = BytesMut::new(); + encode_packet(packet, &mut buf).map_err(|e| format!("Packet encoding failed: {e}"))?; + match maximum_packet_size { + Some(max) if u32::try_from(buf.len()).map_or(true, |len| len > max) => Err(format!( + "{} of {} bytes exceeds the server Maximum Packet Size {max}", + packet.packet_type_name(), + buf.len() + )), + _ => Ok(buf), + } +} + +pub fn write_packet( + state: &Rc>, + packet: &Packet, +) -> std::result::Result<(), String> { + let (writer, maximum_packet_size) = { + let state_ref = state.borrow(); + ( + state_ref.writer.clone(), + state_ref.server.maximum_packet_size, + ) + }; + let buf = encode_checked(packet, maximum_packet_size)?; + let writer = writer.ok_or_else(|| "Not connected".to_string())?; + let result = writer + .borrow_mut() + .write(&buf) + .map_err(|e| format!("Write failed: {e}")); + result +} diff --git a/crates/mqtt5-wasm/src/client/qos.rs b/crates/mqtt5-wasm/src/client/qos.rs index d0f0dd3b..07a354f1 100644 --- a/crates/mqtt5-wasm/src/client/qos.rs +++ b/crates/mqtt5-wasm/src/client/qos.rs @@ -1,12 +1,18 @@ +use mqtt5_protocol::packet::publish::PublishPacket; +use mqtt5_protocol::packet::pubrel::PubRelPacket; +use mqtt5_protocol::packet::Packet; use mqtt5_protocol::QoS; use std::cell::RefCell; use std::rc::Rc; use wasm_bindgen::prelude::*; use wasm_bindgen_futures::JsFuture; +use super::packet::write_packet; use super::sleep_ms; use super::state::ClientState; +const QOS2_CALLBACK_TIMEOUT_MS: f64 = 10_000.0; + pub fn create_ack_promises( state: &Rc>, qos: QoS, @@ -41,32 +47,153 @@ pub fn create_ack_promises( (puback_promise, pubcomp_promise) } -#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)] pub async fn await_ack_promises( puback_promise: Option, pubcomp_promise: Option, ) -> Result<(), JsValue> { - if let Some(promise) = puback_promise { + for promise in [puback_promise, pubcomp_promise].into_iter().flatten() { let result = JsFuture::from(promise).await?; - let reason_code = result.as_f64().unwrap_or(0.0) as u8; - if reason_code >= 0x80 { - return Err(JsValue::from_str(&format!( - "Publish rejected with reason code: {reason_code}" - ))); + match result.as_f64() { + Some(reason_code) if reason_code >= 128.0 => { + return Err(JsValue::from_str(&format!( + "Publish rejected with reason code: {reason_code}" + ))); + } + Some(_) => {} + None => { + return Err(JsValue::from_str(&format!( + "Publish not acknowledged: {}", + result.as_string().unwrap_or_default() + ))); + } } } + Ok(()) +} - if let Some(promise) = pubcomp_promise { - let result = JsFuture::from(promise).await?; - let reason_code = result.as_f64().unwrap_or(0.0) as u8; - if reason_code >= 0x80 { - return Err(JsValue::from_str(&format!( - "Publish rejected with reason code: {reason_code}" - ))); +fn stored_copy(state: &ClientState, publish: &PublishPacket) -> PublishPacket { + let mut stored = publish.clone(); + if let Some(alias) = stored.topic_alias() { + if stored.topic_name.is_empty() { + if let Some(topic) = state.outbound_aliases.get_topic(alias) { + stored.topic_name = topic.to_string(); + } } + stored.properties.remove_topic_alias(); } + stored +} - Ok(()) +pub async fn reserve_flight( + state: &Rc>, + publish: &PublishPacket, +) -> Result { + loop { + { + let mut state_mut = state.borrow_mut(); + if !state_mut.connected { + return Err(JsValue::from_str("Not connected")); + } + if state_mut.send_quota > 0 { + let packet_id = state_mut + .allocate_packet_id() + .ok_or_else(|| JsValue::from_str("No packet identifier available"))?; + state_mut.send_quota -= 1; + let mut stored = stored_copy(&state_mut, publish); + stored.packet_id = Some(packet_id); + state_mut.record_flight(packet_id, stored); + return Ok(packet_id); + } + } + let waiting_state = Rc::clone(state); + let promise = js_sys::Promise::new(&mut move |resolve, _reject| { + waiting_state.borrow_mut().quota_waiters.push_back(resolve); + }); + JsFuture::from(promise).await?; + } +} + +pub fn abandon_flight(state: &Rc>, packet_id: u16) { + let removed = state.borrow_mut().outbound.remove(&packet_id).is_some(); + if removed { + release_quota(state); + } +} + +pub fn release_quota(state: &Rc>) { + let (resend, waiter) = { + let mut state_mut = state.borrow_mut(); + state_mut.send_quota = state_mut + .send_quota + .saturating_add(1) + .min(state_mut.server.receive_maximum); + let mut resend = None; + while let Some(packet_id) = state_mut.pending_resends.pop_front() { + if let Some(flight) = state_mut.outbound.get(&packet_id) { + resend = Some(flight.publish.clone().with_dup(true)); + break; + } + } + match resend { + Some(publish) => { + state_mut.send_quota -= 1; + (Some(publish), None) + } + None => (None, state_mut.quota_waiters.pop_front()), + } + }; + if let Some(publish) = resend { + if let Err(e) = write_packet(state, &Packet::Publish(publish)) { + tracing::warn!(error = %e, "resending PUBLISH failed"); + } + } + if let Some(waiter) = waiter { + if let Err(e) = waiter.call0(&JsValue::NULL) { + tracing::warn!(error = ?e, "send quota waiter failed"); + } + } +} + +pub fn wake_quota_waiters(state: &Rc>) { + let waiters: Vec = state.borrow_mut().quota_waiters.drain(..).collect(); + for waiter in waiters { + if let Err(e) = waiter.call0(&JsValue::NULL) { + tracing::warn!(error = ?e, "send quota waiter failed"); + } + } +} + +pub fn resend_session(state: &Rc>) { + let packets = { + let mut state_mut = state.borrow_mut(); + let mut order: Vec<(u64, u16)> = state_mut + .outbound + .iter() + .map(|(packet_id, flight)| (flight.sequence, *packet_id)) + .collect(); + order.sort_unstable(); + let mut packets = Vec::with_capacity(order.len()); + for (_, packet_id) in order { + let Some(flight) = state_mut.outbound.get(&packet_id) else { + continue; + }; + if flight.released { + packets.push(Packet::PubRel(PubRelPacket::new(packet_id))); + state_mut.send_quota = state_mut.send_quota.saturating_sub(1); + } else if state_mut.send_quota > 0 { + packets.push(Packet::Publish(flight.publish.clone().with_dup(true))); + state_mut.send_quota -= 1; + } else { + state_mut.pending_resends.push_back(packet_id); + } + } + packets + }; + for packet in packets { + if let Err(e) = write_packet(state, &packet) { + tracing::warn!(error = %e, "resending session state failed"); + } + } } pub fn spawn_qos2_cleanup_task(state: Rc>) { @@ -75,53 +202,36 @@ pub fn spawn_qos2_cleanup_task(state: Rc>) { loop { sleep_ms(5000).await; - let connected = match state.try_borrow() { - Ok(state_ref) => { - if state_ref.connection_generation != generation { + let timed_out = match state.try_borrow_mut() { + Ok(mut state_ref) => { + if state_ref.connection_generation != generation || !state_ref.connected { break; } - state_ref.connected + let now = js_sys::Date::now(); + let expired: Vec = state_ref + .pending_pubcomps + .iter() + .filter(|(_, (_, timestamp))| now - timestamp > QOS2_CALLBACK_TIMEOUT_MS) + .map(|(packet_id, _)| *packet_id) + .collect(); + expired + .into_iter() + .filter_map(|packet_id| { + state_ref + .pending_pubcomps + .remove(&packet_id) + .map(|(callback, _)| (packet_id, callback)) + }) + .collect::>() } Err(_) => continue, }; - if !connected { - break; - } - - let now = js_sys::Date::now(); - let timeout_ms = 10000.0; - let cleanup_ms = 30000.0; - - if let Ok(mut state_ref) = state.try_borrow_mut() { - let mut timed_out_pubcomps = Vec::new(); - - for (packet_id, (callback, timestamp)) in &state_ref.pending_pubcomps { - if now - timestamp > timeout_ms { - timed_out_pubcomps.push((*packet_id, callback.clone())); - } - } - - for (packet_id, callback) in timed_out_pubcomps { - state_ref.pending_pubcomps.remove(&packet_id); - web_sys::console::warn_1( - &format!("QoS 2 publish timeout for packet {packet_id}").into(), - ); - let error = JsValue::from_str("Timeout"); - if let Err(e) = callback.call1(&JsValue::NULL, &error) { - web_sys::console::error_1( - &format!("QoS 2 timeout callback error: {e:?}").into(), - ); - } + for (packet_id, callback) in timed_out { + tracing::warn!(packet_id, "QoS 2 publish acknowledgement timed out"); + if let Err(e) = callback.call1(&JsValue::NULL, &JsValue::from_str("Timeout")) { + tracing::warn!(error = ?e, "QoS 2 timeout callback failed"); } - - state_ref - .pending_pubrecs - .retain(|_packet_id, timestamp| now - *timestamp <= cleanup_ms); - - state_ref - .received_qos2 - .retain(|_packet_id, timestamp| now - *timestamp <= cleanup_ms); } } }); diff --git a/crates/mqtt5-wasm/src/client/reader.rs b/crates/mqtt5-wasm/src/client/reader.rs index 75f54523..d676c133 100644 --- a/crates/mqtt5-wasm/src/client/reader.rs +++ b/crates/mqtt5-wasm/src/client/reader.rs @@ -2,40 +2,133 @@ use std::cell::RefCell; use std::rc::Rc; use wasm_bindgen_futures::spawn_local; -use crate::decoder::read_packet; +use crate::decoder::read_frame; use crate::transport::WasmReader; -use mqtt5_protocol::packet::Packet; +use mqtt5_protocol::error::MqttError; +use mqtt5_protocol::packet::{FixedHeader, Packet}; +use mqtt5_protocol::protocol::v5::reason_codes::ReasonCode; -use super::callbacks::handle_connection_lost; -use super::handlers::handle_incoming_packet; +use super::callbacks::{fail_connection, handle_connection_lost}; +use super::handlers::{handle_incoming_packet, Violation}; use super::state::ClientState; +enum ReadOutcome { + Packet(Packet), + Violation(Violation), + Lost(String), +} + +pub fn classify_read_error(error: MqttError) -> Result { + match error { + MqttError::PacketTooLarge { size, max } => Ok(Violation::new( + ReasonCode::PacketTooLarge, + format!("Inbound packet of {size} bytes exceeds Maximum Packet Size {max}"), + )), + MqttError::ProtocolError(message) => Ok(Violation::new(ReasonCode::ProtocolError, message)), + MqttError::ConnectionClosedByPeer + | MqttError::ClientClosed + | MqttError::NotConnected + | MqttError::Io(_) + | MqttError::ConnectionError(_) => Err(format!("Packet read error: {error}")), + other => Ok(Violation::new( + ReasonCode::MalformedPacket, + other.to_string(), + )), + } +} + +pub fn decode_frame( + fixed_header: &FixedHeader, + body: &[u8], + protocol_version: u8, +) -> Result { + let mut cursor = body; + Packet::decode_from_body_with_version( + fixed_header.packet_type, + fixed_header, + &mut cursor, + protocol_version, + ) + .and_then(|packet| match packet { + Packet::Publish(_) | Packet::Subscribe(_) | Packet::SubAck(_) | Packet::Unsubscribe(_) + if !fixed_header.validate_flags() => + { + Err(MqttError::MalformedPacket(format!( + "Invalid fixed header flags 0x{:02X}", + fixed_header.flags + ))) + } + packet => Ok(packet), + }) + .map_err(|e| { + classify_read_error(e) + .unwrap_or_else(|message| Violation::new(ReasonCode::MalformedPacket, message)) + }) +} + +async fn next_packet(state: &Rc>, reader: &mut WasmReader) -> ReadOutcome { + let (maximum_packet_size, protocol_version) = { + let state_ref = state.borrow(); + ( + state_ref.client_limits.maximum_packet_size, + state_ref.protocol_version, + ) + }; + match read_frame(reader, maximum_packet_size).await { + Ok((fixed_header, body)) => match decode_frame(&fixed_header, &body, protocol_version) { + Ok(packet) => ReadOutcome::Packet(packet), + Err(violation) => ReadOutcome::Violation(violation), + }, + Err(error) => match classify_read_error(error) { + Ok(violation) => ReadOutcome::Violation(violation), + Err(message) => ReadOutcome::Lost(message), + }, + } +} + +pub async fn read_connect_response( + state: &Rc>, + reader: &mut WasmReader, +) -> Result { + match next_packet(state, reader).await { + ReadOutcome::Packet(packet) => Ok(packet), + ReadOutcome::Violation(violation) => Err(format!( + "Invalid packet during connect ({:?}): {}", + violation.reason_code, violation.message + )), + ReadOutcome::Lost(message) => Err(message), + } +} + pub fn spawn_packet_reader(state: Rc>, mut reader: WasmReader) { let generation = state.borrow().connection_generation; spawn_local(async move { - let disconnect_reason = loop { - match read_packet(&mut reader).await { - Ok(packet) => { - if state.borrow().connection_generation != generation { - break None; - } - if let Packet::Disconnect(ref disc) = packet { - break Some(format!("Server sent DISCONNECT: {:?}", disc.reason_code)); + loop { + let outcome = next_packet(&state, &mut reader).await; + if state.borrow().connection_generation != generation { + return; + } + match outcome { + ReadOutcome::Packet(Packet::Disconnect(disconnect)) => { + let reason = format!("Server sent DISCONNECT: {:?}", disconnect.reason_code); + handle_connection_lost(&state, &reason); + return; + } + ReadOutcome::Packet(packet) => { + if let Err(violation) = handle_incoming_packet(&state, packet) { + fail_connection(&state, &violation); + return; } - handle_incoming_packet(&state, packet); } - Err(e) => { - break Some(format!("Packet read error: {e}")); + ReadOutcome::Violation(violation) => { + fail_connection(&state, &violation); + return; + } + ReadOutcome::Lost(message) => { + handle_connection_lost(&state, &message); + return; } } - }; - - if state.borrow().connection_generation != generation { - return; - } - - if let Some(reason) = disconnect_reason { - handle_connection_lost(&state, &reason); } }); } diff --git a/crates/mqtt5-wasm/src/client/reconnect.rs b/crates/mqtt5-wasm/src/client/reconnect.rs index 72371d92..19cc7d6f 100644 --- a/crates/mqtt5-wasm/src/client/reconnect.rs +++ b/crates/mqtt5-wasm/src/client/reconnect.rs @@ -1,22 +1,15 @@ -use bytes::BytesMut; use mqtt5_protocol::packet::connect::ConnectPacket; -use mqtt5_protocol::packet::Packet; use mqtt5_protocol::protocol::v5::properties::Properties; use mqtt5_protocol::u128_to_u32_saturating; -use mqtt5_protocol::Transport; use std::cell::RefCell; use std::rc::Rc; use wasm_bindgen_futures::spawn_local; -use crate::decoder::read_packet; use crate::transport::WasmTransportType; use super::callbacks::{trigger_reconnect_failed_callback, trigger_reconnecting_callback}; +use super::connection::establish; use super::connectivity::{is_browser_online, wait_for_online}; -use super::keepalive::spawn_keepalive_task; -use super::packet::encode_packet; -use super::qos::spawn_qos2_cleanup_task; -use super::reader::spawn_packet_reader; use super::sleep_ms; use super::state::{ClientState, StoredConnectOptions}; @@ -34,13 +27,10 @@ pub fn spawn_reconnection_task(state: Rc>) { let attempt = state_ref.reconnect_attempt; let base_delay = state_ref.reconnect_config.calculate_delay(attempt); let base_delay_ms = u128_to_u32_saturating(base_delay.as_millis()); - let jitter_f64 = js_sys::Math::random() * f64::from(base_delay_ms / 4); - let jitter = if jitter_f64 >= f64::from(u32::MAX) { - u32::MAX - } else { - #[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)] - let result = jitter_f64 as u32; - result + let jitter_range = base_delay_ms / 4; + let jitter = match jitter_range { + 0 => 0, + range => getrandom::u32().unwrap_or(0) % range, }; let delay_ms = base_delay_ms.saturating_add(jitter); let should_continue = state_ref.reconnect_config.should_retry(attempt); @@ -107,34 +97,16 @@ async fn try_all_brokers( state_ref.reconnect_attempt = 0; state_ref.current_broker_index = idx; } - if idx == 0 { - web_sys::console::log_1( - &format!("Reconnected to primary after {0} attempt(s)", attempt + 1).into(), - ); - } else { - web_sys::console::log_1( - &format!( - "Reconnected to backup {idx} after {0} attempt(s)", - attempt + 1 - ) - .into(), - ); - } + tracing::info!(broker_index = idx, attempts = attempt + 1, "reconnected"); return true; } Err(e) => { - if idx == 0 { - web_sys::console::warn_1(&format!("Primary connection failed: {e}").into()); - } else { - web_sys::console::warn_1( - &format!("Backup {idx} connection failed: {e}").into(), - ); - } + tracing::warn!(broker_index = idx, error = %e, "reconnection attempt failed"); } } } - web_sys::console::warn_1(&format!("All brokers failed on attempt {0}", attempt + 1).into()); + tracing::warn!(attempt = attempt + 1, "all brokers failed"); false } @@ -143,28 +115,15 @@ async fn attempt_reconnect( url: &str, options: &StoredConnectOptions, ) -> Result<(), String> { - let mut transport = WasmTransportType::WebSocket( + let transport = WasmTransportType::WebSocket( crate::transport::websocket::WasmWebSocketTransport::new(url), ); - - transport - .connect() - .await - .map_err(|e| format!("Transport connection failed: {e}"))?; - let client_id = state.borrow().client_id.clone(); - - { - let mut state_mut = state.borrow_mut(); - state_mut.keep_alive = options.keep_alive; - state_mut.protocol_version = options.protocol_version; - #[cfg(feature = "codec")] - { - state_mut.codec_registry.clone_from(&options.codec_registry); - } - } - - let properties = build_properties_from_stored(options); + let properties = if options.protocol_version == 5 { + build_properties_from_stored(options) + } else { + Properties::default() + }; let connect_packet = ConnectPacket { protocol_version: options.protocol_version, @@ -178,134 +137,49 @@ async fn attempt_reconnect( will_properties: Properties::default(), }; - let packet = Packet::Connect(Box::new(connect_packet)); - let mut buf = BytesMut::new(); - encode_packet(&packet, &mut buf).map_err(|e| format!("Packet encoding failed: {e}"))?; - - transport - .write(&buf) - .await - .map_err(|e| format!("Write failed: {e}"))?; - - if let Some(method) = &options.authentication_method { - state.borrow_mut().auth_method = Some(method.clone()); - } - - let (reader, writer) = transport - .into_split() - .map_err(|e| format!("Transport split failed: {e}"))?; - - let mut reader = reader; - let packet = read_packet(&mut reader) + establish(state, transport, connect_packet, options) .await - .map_err(|e| format!("Packet read failed: {e}"))?; - - match packet { - Packet::ConnAck(connack) => { - let reason_code = connack.reason_code as u8; - if reason_code != 0 { - return Err(format!("CONNACK error: {reason_code}")); - } - - let writer_rc = Rc::new(RefCell::new(writer)); - { - let mut state_mut = state.borrow_mut(); - state_mut.connected = true; - state_mut.connection_generation = state_mut.connection_generation.wrapping_add(1); - state_mut.writer = Some(Rc::clone(&writer_rc)); - } - - spawn_packet_reader(Rc::clone(state), reader); - spawn_keepalive_task(Rc::clone(state)); - spawn_qos2_cleanup_task(Rc::clone(state)); - - let callback = state.borrow().on_connect.clone(); - if let Some(callback) = callback { - let reason_code_js = - wasm_bindgen::JsValue::from_f64(f64::from(connack.reason_code as u8)); - let session_present_js = wasm_bindgen::JsValue::from_bool(connack.session_present); - let _ = callback.call2( - &wasm_bindgen::JsValue::NULL, - &reason_code_js, - &session_present_js, - ); - } - - Ok(()) - } - _ => Err(format!("Expected CONNACK, received: {packet:?}")), - } + .map(|_| ()) + .map_err(|failure| failure.to_string()) } pub fn build_properties_from_stored(options: &StoredConnectOptions) -> Properties { let mut properties = Properties::default(); if let Some(interval) = options.session_expiry_interval { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::SessionExpiryInterval, - mqtt5_protocol::protocol::v5::properties::PropertyValue::FourByteInteger(interval), - ); + properties.set_session_expiry_interval(interval); } if let Some(max) = options.receive_maximum { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::ReceiveMaximum, - mqtt5_protocol::protocol::v5::properties::PropertyValue::TwoByteInteger(max), - ); + properties.set_receive_maximum(max); } if let Some(size) = options.maximum_packet_size { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::MaximumPacketSize, - mqtt5_protocol::protocol::v5::properties::PropertyValue::FourByteInteger(size), - ); + properties.set_maximum_packet_size(size); } if let Some(max) = options.topic_alias_maximum { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::TopicAliasMaximum, - mqtt5_protocol::protocol::v5::properties::PropertyValue::TwoByteInteger(max), - ); + properties.set_topic_alias_maximum(max); } if let Some(req) = options.request_response_information { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::RequestResponseInformation, - mqtt5_protocol::protocol::v5::properties::PropertyValue::Byte(u8::from(req)), - ); + properties.set_request_response_information(req); } if let Some(req) = options.request_problem_information { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::RequestProblemInformation, - mqtt5_protocol::protocol::v5::properties::PropertyValue::Byte(u8::from(req)), - ); + properties.set_request_problem_information(req); } if let Some(method) = &options.authentication_method { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::AuthenticationMethod, - mqtt5_protocol::protocol::v5::properties::PropertyValue::Utf8String(method.clone()), - ); + properties.set_authentication_method(method.clone()); } if let Some(data) = &options.authentication_data { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::AuthenticationData, - mqtt5_protocol::protocol::v5::properties::PropertyValue::BinaryData( - data.clone().into(), - ), - ); + properties.set_authentication_data(data.clone().into()); } for (key, value) in &options.user_properties { - let _ = properties.add( - mqtt5_protocol::protocol::v5::properties::PropertyId::UserProperty, - mqtt5_protocol::protocol::v5::properties::PropertyValue::Utf8StringPair( - key.clone(), - value.clone(), - ), - ); + properties.add_user_property(key.clone(), value.clone()); } properties diff --git a/crates/mqtt5-wasm/src/client/state.rs b/crates/mqtt5-wasm/src/client/state.rs index 07647acc..3408ece7 100644 --- a/crates/mqtt5-wasm/src/client/state.rs +++ b/crates/mqtt5-wasm/src/client/state.rs @@ -1,14 +1,125 @@ use crate::config::WasmConnectOptions; use crate::transport::WasmWriter; use mqtt5_protocol::connection::ReconnectConfig; +use mqtt5_protocol::packet::connack::ConnAckPacket; +use mqtt5_protocol::packet::publish::PublishPacket; use mqtt5_protocol::packet_id::PacketIdGenerator; +use mqtt5_protocol::protocol::v5::properties::{PropertyId, PropertyValue}; +use mqtt5_protocol::session::TopicAliasManager; +use mqtt5_protocol::QoS; use std::cell::RefCell; -use std::collections::HashMap; +use std::collections::{HashMap, HashSet, VecDeque}; use std::rc::Rc; #[cfg(feature = "codec")] use crate::codec::WasmCodecRegistry; +const DEFAULT_RECEIVE_MAXIMUM: u16 = u16::MAX; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ServerLimits { + pub receive_maximum: u16, + pub maximum_qos: QoS, + pub retain_available: bool, + pub maximum_packet_size: Option, + pub topic_alias_maximum: u16, + pub subscriptions: SubscriptionFeatures, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct SubscriptionFeatures { + pub wildcards: bool, + pub shared: bool, + pub identifiers: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SessionState { + Absent, + Held, +} + +impl Default for ServerLimits { + fn default() -> Self { + Self { + receive_maximum: DEFAULT_RECEIVE_MAXIMUM, + maximum_qos: QoS::ExactlyOnce, + retain_available: true, + maximum_packet_size: None, + topic_alias_maximum: 0, + subscriptions: SubscriptionFeatures { + wildcards: true, + shared: true, + identifiers: true, + }, + } + } +} + +impl ServerLimits { + pub fn from_connack(connack: &ConnAckPacket) -> Self { + let properties = &connack.properties; + let available = + |id: PropertyId| !matches!(properties.get(id), Some(PropertyValue::Byte(0))); + Self { + receive_maximum: connack + .receive_maximum() + .filter(|max| *max > 0) + .unwrap_or(DEFAULT_RECEIVE_MAXIMUM), + maximum_qos: match properties.get_maximum_qos() { + Some(0) => QoS::AtMostOnce, + Some(1) => QoS::AtLeastOnce, + _ => QoS::ExactlyOnce, + }, + retain_available: available(PropertyId::RetainAvailable), + maximum_packet_size: connack.maximum_packet_size(), + topic_alias_maximum: connack.topic_alias_maximum().unwrap_or(0), + subscriptions: SubscriptionFeatures { + wildcards: available(PropertyId::WildcardSubscriptionAvailable), + shared: available(PropertyId::SharedSubscriptionAvailable), + identifiers: available(PropertyId::SubscriptionIdentifierAvailable), + }, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct ClientLimits { + pub receive_maximum: u16, + pub maximum_packet_size: u32, + pub topic_alias_maximum: u16, + pub request_problem_information: bool, +} + +impl Default for ClientLimits { + fn default() -> Self { + Self::from(&StoredConnectOptions::from(&WasmConnectOptions::default())) + } +} + +impl From<&StoredConnectOptions> for ClientLimits { + fn from(options: &StoredConnectOptions) -> Self { + Self { + receive_maximum: options + .receive_maximum + .filter(|max| *max > 0) + .unwrap_or(DEFAULT_RECEIVE_MAXIMUM), + maximum_packet_size: options + .maximum_packet_size + .filter(|size| *size > 0) + .unwrap_or(mqtt5_protocol::constants::limits::MAX_PACKET_SIZE), + topic_alias_maximum: options.topic_alias_maximum.unwrap_or(0), + request_problem_information: options.request_problem_information.unwrap_or(true), + } + } +} + +pub struct OutboundFlight { + pub sequence: u64, + pub publish: PublishPacket, + pub released: bool, +} + pub struct ClientState { pub client_id: String, pub writer: Option>>, @@ -17,11 +128,22 @@ pub struct ClientState { pub protocol_version: u8, pub subscriptions: HashMap, pub rust_subscriptions: HashMap, - pub pending_subacks: HashMap, + pub pending_subacks: HashMap>, + pub pending_unsubacks: HashSet, pub pending_pubacks: HashMap, pub pending_pubcomps: HashMap, - pub pending_pubrecs: HashMap, - pub received_qos2: HashMap, + pub outbound: HashMap, + pub next_flight_sequence: u64, + pub send_quota: u16, + pub quota_waiters: VecDeque, + pub pending_resends: VecDeque, + pub awaiting_pubrel: HashSet, + pub server: ServerLimits, + pub client_limits: ClientLimits, + pub outbound_aliases: TopicAliasManager, + pub inbound_aliases: TopicAliasManager, + pub session: SessionState, + pub session_expiry_interval: u32, pub keep_alive: u16, pub last_ping_sent: Option, pub last_pong_received: Option, @@ -58,10 +180,21 @@ impl ClientState { subscriptions: HashMap::new(), rust_subscriptions: HashMap::new(), pending_subacks: HashMap::new(), + pending_unsubacks: HashSet::new(), pending_pubacks: HashMap::new(), pending_pubcomps: HashMap::new(), - pending_pubrecs: HashMap::new(), - received_qos2: HashMap::new(), + outbound: HashMap::new(), + next_flight_sequence: 0, + send_quota: DEFAULT_RECEIVE_MAXIMUM, + quota_waiters: VecDeque::new(), + pending_resends: VecDeque::new(), + awaiting_pubrel: HashSet::new(), + server: ServerLimits::default(), + client_limits: ClientLimits::default(), + outbound_aliases: TopicAliasManager::new(0), + inbound_aliases: TopicAliasManager::new(0), + session: SessionState::Absent, + session_expiry_interval: 0, keep_alive: 60, last_ping_sent: None, last_pong_received: None, @@ -87,11 +220,84 @@ impl ClientState { codec_registry: None, } } + + pub fn packet_id_in_use(&self, packet_id: u16) -> bool { + self.outbound.contains_key(&packet_id) + || self.pending_subacks.contains_key(&packet_id) + || self.pending_unsubacks.contains(&packet_id) + } + + pub fn allocate_packet_id(&self) -> Option { + (0..=u16::MAX).find_map(|_| { + let packet_id = self.packet_id.next(); + (!self.packet_id_in_use(packet_id)).then_some(packet_id) + }) + } + + pub fn record_flight(&mut self, packet_id: u16, publish: PublishPacket) { + let sequence = self.next_flight_sequence; + self.next_flight_sequence = sequence.wrapping_add(1); + self.outbound.insert( + packet_id, + OutboundFlight { + sequence, + publish, + released: false, + }, + ); + } + + pub fn discard_session(&mut self) -> Vec { + self.outbound.clear(); + self.pending_resends.clear(); + self.awaiting_pubrel.clear(); + self.session = SessionState::Absent; + self.pending_pubacks + .drain() + .map(|(_, callback)| callback) + .chain( + self.pending_pubcomps + .drain() + .map(|(_, (callback, _))| callback), + ) + .collect() + } + + pub fn apply_connect_options(&mut self, options: &StoredConnectOptions) { + self.keep_alive = options.keep_alive; + self.protocol_version = options.protocol_version; + self.session_expiry_interval = options.session_expiry_interval.unwrap_or(0); + self.client_limits = ClientLimits::from(options); + self.auth_method.clone_from(&options.authentication_method); + #[cfg(feature = "codec")] + { + self.codec_registry.clone_from(&options.codec_registry); + } + } + + pub fn apply_connack(&mut self, connack: &ConnAckPacket) { + self.server = ServerLimits::from_connack(connack); + self.send_quota = self.server.receive_maximum; + self.outbound_aliases = TopicAliasManager::new(self.server.topic_alias_maximum); + self.inbound_aliases = TopicAliasManager::new(self.client_limits.topic_alias_maximum); + self.pending_resends.clear(); + if let Some(keep_alive) = connack.properties.get_server_keep_alive() { + self.keep_alive = keep_alive; + } + if let Some(PropertyValue::Utf8String(assigned)) = + connack.properties.get(PropertyId::AssignedClientIdentifier) + { + self.client_id.clone_from(assigned); + } + self.last_ping_sent = None; + self.last_pong_received = None; + } } #[derive(Clone)] pub struct StoredConnectOptions { pub keep_alive: u16, + pub resume_existing_session: bool, pub username: Option, pub password: Option>, pub session_expiry_interval: Option, @@ -113,6 +319,7 @@ impl From<&WasmConnectOptions> for StoredConnectOptions { fn from(opts: &WasmConnectOptions) -> Self { Self { keep_alive: opts.keep_alive, + resume_existing_session: opts.resume_existing_session, username: opts.username.clone(), password: opts.password.clone(), session_expiry_interval: opts.session_expiry_interval, diff --git a/crates/mqtt5-wasm/src/config.rs b/crates/mqtt5-wasm/src/config.rs index b56b3b3b..23287c30 100644 --- a/crates/mqtt5-wasm/src/config.rs +++ b/crates/mqtt5-wasm/src/config.rs @@ -26,10 +26,9 @@ pub struct WasmReconnectOptions { } #[wasm_bindgen(js_class = "ReconnectOptions")] -#[allow(non_snake_case)] impl WasmReconnectOptions { #[wasm_bindgen(constructor)] - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new() -> Self { Self { enabled: true, @@ -59,47 +58,47 @@ impl WasmReconnectOptions { self.enabled = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = initialDelayMs)] #[must_use] - pub fn initialDelayMs(&self) -> u32 { + pub fn initial_delay_ms(&self) -> u32 { self.initial_delay_ms } - #[wasm_bindgen(setter)] - pub fn set_initialDelayMs(&mut self, value: u32) { + #[wasm_bindgen(setter = initialDelayMs)] + pub fn set_initial_delay_ms(&mut self, value: u32) { self.initial_delay_ms = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = maxDelayMs)] #[must_use] - pub fn maxDelayMs(&self) -> u32 { + pub fn max_delay_ms(&self) -> u32 { self.max_delay_ms } - #[wasm_bindgen(setter)] - pub fn set_maxDelayMs(&mut self, value: u32) { + #[wasm_bindgen(setter = maxDelayMs)] + pub fn set_max_delay_ms(&mut self, value: u32) { self.max_delay_ms = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = backoffFactor)] #[must_use] - pub fn backoffFactor(&self) -> f64 { + pub fn backoff_factor(&self) -> f64 { self.backoff_factor } - #[wasm_bindgen(setter)] - pub fn set_backoffFactor(&mut self, value: f64) { + #[wasm_bindgen(setter = backoffFactor)] + pub fn set_backoff_factor(&mut self, value: f64) { self.backoff_factor = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = maxAttempts)] #[must_use] - pub fn maxAttempts(&self) -> Option { + pub fn max_attempts(&self) -> Option { self.max_attempts } - #[wasm_bindgen(setter)] - pub fn set_maxAttempts(&mut self, value: Option) { + #[wasm_bindgen(setter = maxAttempts)] + pub fn set_max_attempts(&mut self, value: Option) { self.max_attempts = value; } @@ -139,6 +138,7 @@ impl Clone for WasmReconnectOptions { pub struct WasmConnectOptions { pub(crate) keep_alive: u16, pub(crate) clean_start: bool, + pub(crate) resume_existing_session: bool, pub(crate) username: Option, pub(crate) password: Option>, pub(crate) will: Option, @@ -158,14 +158,14 @@ pub struct WasmConnectOptions { } #[wasm_bindgen(js_class = "ConnectOptions")] -#[allow(non_snake_case)] impl WasmConnectOptions { #[wasm_bindgen(constructor)] - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new() -> Self { Self { keep_alive: 60, clean_start: true, + resume_existing_session: false, username: None, password: None, will: None, @@ -185,28 +185,49 @@ impl WasmConnectOptions { } } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = keepAlive)] #[must_use] - pub fn keepAlive(&self) -> u16 { + pub fn keep_alive(&self) -> u16 { self.keep_alive } - #[wasm_bindgen(setter)] - pub fn set_keepAlive(&mut self, value: u16) { + #[wasm_bindgen(setter = keepAlive)] + pub fn set_keep_alive(&mut self, value: u16) { self.keep_alive = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = cleanStart)] #[must_use] - pub fn cleanStart(&self) -> bool { + pub fn clean_start(&self) -> bool { self.clean_start } - #[wasm_bindgen(setter)] - pub fn set_cleanStart(&mut self, value: bool) { + #[wasm_bindgen(setter = cleanStart)] + pub fn set_clean_start(&mut self, value: bool) { self.clean_start = value; } + /// Accept `Session Present = 1` on a client that holds no local session state. + /// + /// By default a client that has not yet established a session in this instance + /// rejects a CONNACK with Session Present set to 1, sends DISCONNECT 0x82 and closes + /// the connection (MQTT-3.2.2-4). Set this to `true` (together with + /// `cleanStart = false`) to deliberately resume a session the broker still holds, + /// for example after a page reload or crash. The client has nothing to resend in + /// that case; the broker resumes delivery of the messages it holds. A CONNACK with + /// Session Present set to 1 in reply to `cleanStart = true` is always rejected. + #[wasm_bindgen(getter = resumeExistingSession)] + #[must_use] + pub fn resume_existing_session(&self) -> bool { + self.resume_existing_session + } + + /// Sets [`Self::resume_existing_session`]. + #[wasm_bindgen(setter = resumeExistingSession)] + pub fn set_resume_existing_session(&mut self, value: bool) { + self.resume_existing_session = value; + } + #[wasm_bindgen(getter)] #[must_use] pub fn username(&self) -> Option { @@ -233,96 +254,96 @@ impl WasmConnectOptions { self.will = None; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = sessionExpiryInterval)] #[must_use] - pub fn sessionExpiryInterval(&self) -> Option { + pub fn session_expiry_interval(&self) -> Option { self.session_expiry_interval } - #[wasm_bindgen(setter)] - pub fn set_sessionExpiryInterval(&mut self, value: Option) { + #[wasm_bindgen(setter = sessionExpiryInterval)] + pub fn set_session_expiry_interval(&mut self, value: Option) { self.session_expiry_interval = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = receiveMaximum)] #[must_use] - pub fn receiveMaximum(&self) -> Option { + pub fn receive_maximum(&self) -> Option { self.receive_maximum } - #[wasm_bindgen(setter)] - pub fn set_receiveMaximum(&mut self, value: Option) { + #[wasm_bindgen(setter = receiveMaximum)] + pub fn set_receive_maximum(&mut self, value: Option) { self.receive_maximum = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = maximumPacketSize)] #[must_use] - pub fn maximumPacketSize(&self) -> Option { + pub fn maximum_packet_size(&self) -> Option { self.maximum_packet_size } - #[wasm_bindgen(setter)] - pub fn set_maximumPacketSize(&mut self, value: Option) { + #[wasm_bindgen(setter = maximumPacketSize)] + pub fn set_maximum_packet_size(&mut self, value: Option) { self.maximum_packet_size = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = topicAliasMaximum)] #[must_use] - pub fn topicAliasMaximum(&self) -> Option { + pub fn topic_alias_maximum(&self) -> Option { self.topic_alias_maximum } - #[wasm_bindgen(setter)] - pub fn set_topicAliasMaximum(&mut self, value: Option) { + #[wasm_bindgen(setter = topicAliasMaximum)] + pub fn set_topic_alias_maximum(&mut self, value: Option) { self.topic_alias_maximum = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = requestResponseInformation)] #[must_use] - pub fn requestResponseInformation(&self) -> Option { + pub fn request_response_information(&self) -> Option { self.request_response_information } - #[wasm_bindgen(setter)] - pub fn set_requestResponseInformation(&mut self, value: Option) { + #[wasm_bindgen(setter = requestResponseInformation)] + pub fn set_request_response_information(&mut self, value: Option) { self.request_response_information = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = requestProblemInformation)] #[must_use] - pub fn requestProblemInformation(&self) -> Option { + pub fn request_problem_information(&self) -> Option { self.request_problem_information } - #[wasm_bindgen(setter)] - pub fn set_requestProblemInformation(&mut self, value: Option) { + #[wasm_bindgen(setter = requestProblemInformation)] + pub fn set_request_problem_information(&mut self, value: Option) { self.request_problem_information = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = authenticationMethod)] #[must_use] - pub fn authenticationMethod(&self) -> Option { + pub fn authentication_method(&self) -> Option { self.authentication_method.clone() } - #[wasm_bindgen(setter)] - pub fn set_authenticationMethod(&mut self, value: Option) { + #[wasm_bindgen(setter = authenticationMethod)] + pub fn set_authentication_method(&mut self, value: Option) { self.authentication_method = value; } - #[wasm_bindgen(setter)] - pub fn set_authenticationData(&mut self, value: &[u8]) { + #[wasm_bindgen(setter = authenticationData)] + pub fn set_authentication_data(&mut self, value: &[u8]) { self.authentication_data = Some(value.to_vec()); } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = protocolVersion)] #[must_use] - pub fn protocolVersion(&self) -> u8 { + pub fn protocol_version(&self) -> u8 { self.protocol_version } - #[wasm_bindgen(setter)] - pub fn set_protocolVersion(&mut self, value: u8) { + #[wasm_bindgen(setter = protocolVersion)] + pub fn set_protocol_version(&mut self, value: u8) { if value == 4 || value == 5 { self.protocol_version = value; } else { @@ -333,24 +354,29 @@ impl WasmConnectOptions { } } - pub fn addUserProperty(&mut self, key: String, value: String) { + #[wasm_bindgen(js_name = addUserProperty)] + pub fn add_user_property(&mut self, key: String, value: String) { self.user_properties.push((key, value)); } - pub fn clearUserProperties(&mut self) { + #[wasm_bindgen(js_name = clearUserProperties)] + pub fn clear_user_properties(&mut self) { self.user_properties.clear(); } - pub fn addBackupUrl(&mut self, url: String) { + #[wasm_bindgen(js_name = addBackupUrl)] + pub fn add_backup_url(&mut self, url: String) { self.backup_urls.push(url); } - pub fn clearBackupUrls(&mut self) { + #[wasm_bindgen(js_name = clearBackupUrls)] + pub fn clear_backup_urls(&mut self) { self.backup_urls.clear(); } #[must_use] - pub fn getBackupUrls(&self) -> Vec { + #[wasm_bindgen(js_name = getBackupUrls)] + pub fn get_backup_urls(&self) -> Vec { self.backup_urls.clone() } @@ -465,10 +491,9 @@ pub struct WasmPublishOptions { } #[wasm_bindgen(js_class = "PublishOptions")] -#[allow(non_snake_case)] impl WasmPublishOptions { #[wasm_bindgen(constructor)] - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new() -> Self { Self { qos: 0, @@ -510,71 +535,73 @@ impl WasmPublishOptions { self.retain = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = payloadFormatIndicator)] #[must_use] - pub fn payloadFormatIndicator(&self) -> Option { + pub fn payload_format_indicator(&self) -> Option { self.payload_format_indicator } - #[wasm_bindgen(setter)] - pub fn set_payloadFormatIndicator(&mut self, value: Option) { + #[wasm_bindgen(setter = payloadFormatIndicator)] + pub fn set_payload_format_indicator(&mut self, value: Option) { self.payload_format_indicator = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = messageExpiryInterval)] #[must_use] - pub fn messageExpiryInterval(&self) -> Option { + pub fn message_expiry_interval(&self) -> Option { self.message_expiry_interval } - #[wasm_bindgen(setter)] - pub fn set_messageExpiryInterval(&mut self, value: Option) { + #[wasm_bindgen(setter = messageExpiryInterval)] + pub fn set_message_expiry_interval(&mut self, value: Option) { self.message_expiry_interval = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = topicAlias)] #[must_use] - pub fn topicAlias(&self) -> Option { + pub fn topic_alias(&self) -> Option { self.topic_alias } - #[wasm_bindgen(setter)] - pub fn set_topicAlias(&mut self, value: Option) { + #[wasm_bindgen(setter = topicAlias)] + pub fn set_topic_alias(&mut self, value: Option) { self.topic_alias = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = responseTopic)] #[must_use] - pub fn responseTopic(&self) -> Option { + pub fn response_topic(&self) -> Option { self.response_topic.clone() } - #[wasm_bindgen(setter)] - pub fn set_responseTopic(&mut self, value: Option) { + #[wasm_bindgen(setter = responseTopic)] + pub fn set_response_topic(&mut self, value: Option) { self.response_topic = value; } - #[wasm_bindgen(setter)] - pub fn set_correlationData(&mut self, value: &[u8]) { + #[wasm_bindgen(setter = correlationData)] + pub fn set_correlation_data(&mut self, value: &[u8]) { self.correlation_data = Some(value.to_vec()); } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = contentType)] #[must_use] - pub fn contentType(&self) -> Option { + pub fn content_type(&self) -> Option { self.content_type.clone() } - #[wasm_bindgen(setter)] - pub fn set_contentType(&mut self, value: Option) { + #[wasm_bindgen(setter = contentType)] + pub fn set_content_type(&mut self, value: Option) { self.content_type = value; } - pub fn addUserProperty(&mut self, key: String, value: String) { + #[wasm_bindgen(js_name = addUserProperty)] + pub fn add_user_property(&mut self, key: String, value: String) { self.user_properties.push((key, value)); } - pub fn clearUserProperties(&mut self) { + #[wasm_bindgen(js_name = clearUserProperties)] + pub fn clear_user_properties(&mut self) { self.user_properties.clear(); } @@ -690,10 +717,9 @@ pub struct WasmSubscribeOptions { } #[wasm_bindgen(js_class = "SubscribeOptions")] -#[allow(non_snake_case)] impl WasmSubscribeOptions { #[wasm_bindgen(constructor)] - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new() -> Self { Self { qos: 0, @@ -720,36 +746,36 @@ impl WasmSubscribeOptions { } } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = noLocal)] #[must_use] - pub fn noLocal(&self) -> bool { + pub fn no_local(&self) -> bool { self.no_local } - #[wasm_bindgen(setter)] - pub fn set_noLocal(&mut self, value: bool) { + #[wasm_bindgen(setter = noLocal)] + pub fn set_no_local(&mut self, value: bool) { self.no_local = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = retainAsPublished)] #[must_use] - pub fn retainAsPublished(&self) -> bool { + pub fn retain_as_published(&self) -> bool { self.retain_as_published } - #[wasm_bindgen(setter)] - pub fn set_retainAsPublished(&mut self, value: bool) { + #[wasm_bindgen(setter = retainAsPublished)] + pub fn set_retain_as_published(&mut self, value: bool) { self.retain_as_published = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = retainHandling)] #[must_use] - pub fn retainHandling(&self) -> u8 { + pub fn retain_handling(&self) -> u8 { self.retain_handling } - #[wasm_bindgen(setter)] - pub fn set_retainHandling(&mut self, value: u8) { + #[wasm_bindgen(setter = retainHandling)] + pub fn set_retain_handling(&mut self, value: u8) { if value > 2 { web_sys::console::warn_1(&"Retain handling must be 0, 1, or 2. Using 0.".into()); self.retain_handling = 0; @@ -758,14 +784,14 @@ impl WasmSubscribeOptions { } } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = subscriptionIdentifier)] #[must_use] - pub fn subscriptionIdentifier(&self) -> Option { + pub fn subscription_identifier(&self) -> Option { self.subscription_identifier } - #[wasm_bindgen(setter)] - pub fn set_subscriptionIdentifier(&mut self, value: Option) { + #[wasm_bindgen(setter = subscriptionIdentifier)] + pub fn set_subscription_identifier(&mut self, value: Option) { self.subscription_identifier = value; } @@ -797,10 +823,9 @@ pub struct WasmWillMessage { } #[wasm_bindgen(js_class = "WillMessage")] -#[allow(non_snake_case)] impl WasmWillMessage { #[wasm_bindgen(constructor)] - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new(topic: String, payload: Vec) -> Self { Self { topic, @@ -852,47 +877,47 @@ impl WasmWillMessage { self.retain = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = willDelayInterval)] #[must_use] - pub fn willDelayInterval(&self) -> Option { + pub fn will_delay_interval(&self) -> Option { self.will_delay_interval } - #[wasm_bindgen(setter)] - pub fn set_willDelayInterval(&mut self, value: Option) { + #[wasm_bindgen(setter = willDelayInterval)] + pub fn set_will_delay_interval(&mut self, value: Option) { self.will_delay_interval = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = messageExpiryInterval)] #[must_use] - pub fn messageExpiryInterval(&self) -> Option { + pub fn message_expiry_interval(&self) -> Option { self.message_expiry_interval } - #[wasm_bindgen(setter)] - pub fn set_messageExpiryInterval(&mut self, value: Option) { + #[wasm_bindgen(setter = messageExpiryInterval)] + pub fn set_message_expiry_interval(&mut self, value: Option) { self.message_expiry_interval = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = contentType)] #[must_use] - pub fn contentType(&self) -> Option { + pub fn content_type(&self) -> Option { self.content_type.clone() } - #[wasm_bindgen(setter)] - pub fn set_contentType(&mut self, value: Option) { + #[wasm_bindgen(setter = contentType)] + pub fn set_content_type(&mut self, value: Option) { self.content_type = value; } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = responseTopic)] #[must_use] - pub fn responseTopic(&self) -> Option { + pub fn response_topic(&self) -> Option { self.response_topic.clone() } - #[wasm_bindgen(setter)] - pub fn set_responseTopic(&mut self, value: Option) { + #[wasm_bindgen(setter = responseTopic)] + pub fn set_response_topic(&mut self, value: Option) { self.response_topic = value; } @@ -932,46 +957,46 @@ pub struct WasmMessageProperties { } #[wasm_bindgen(js_class = "MessageProperties")] -#[allow(non_snake_case)] impl WasmMessageProperties { - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = responseTopic)] #[must_use] - pub fn responseTopic(&self) -> Option { + pub fn response_topic(&self) -> Option { self.response_topic.clone() } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = correlationData)] #[must_use] - pub fn correlationData(&self) -> Option> { + pub fn correlation_data(&self) -> Option> { self.correlation_data.clone() } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = contentType)] #[must_use] - pub fn contentType(&self) -> Option { + pub fn content_type(&self) -> Option { self.content_type.clone() } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = payloadFormatIndicator)] #[must_use] - pub fn payloadFormatIndicator(&self) -> Option { + pub fn payload_format_indicator(&self) -> Option { self.payload_format_indicator } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = messageExpiryInterval)] #[must_use] - pub fn messageExpiryInterval(&self) -> Option { + pub fn message_expiry_interval(&self) -> Option { self.message_expiry_interval } - #[wasm_bindgen(getter)] + #[wasm_bindgen(getter = subscriptionIdentifiers)] #[must_use] - pub fn subscriptionIdentifiers(&self) -> Vec { + pub fn subscription_identifiers(&self) -> Vec { self.subscription_identifiers.clone() } #[must_use] - pub fn getUserProperties(&self) -> js_sys::Array { + #[wasm_bindgen(js_name = getUserProperties)] + pub fn get_user_properties(&self) -> js_sys::Array { let arr = js_sys::Array::new(); for (key, value) in &self.user_properties { let pair = js_sys::Array::new(); diff --git a/crates/mqtt5-wasm/src/decoder.rs b/crates/mqtt5-wasm/src/decoder.rs index 0b0c295c..7b411103 100644 --- a/crates/mqtt5-wasm/src/decoder.rs +++ b/crates/mqtt5-wasm/src/decoder.rs @@ -1,67 +1,77 @@ use crate::transport::WasmReader; -use bytes::Buf; use mqtt5_protocol::constants::limits::MAX_PACKET_SIZE; use mqtt5_protocol::error::{MqttError, Result}; -use mqtt5_protocol::packet::{FixedHeader, Packet}; +use mqtt5_protocol::packet::{FixedHeader, Packet, PacketType}; + +const MAX_REMAINING_LENGTH_BYTES: usize = 4; /// # Errors /// Returns an error if the connection is closed or packet decoding fails. pub async fn read_packet(reader: &mut WasmReader) -> Result { - let mut header_buf = vec![0u8; 5]; - let n = reader.read(&mut header_buf).await?; - - if n == 0 { - return Err(MqttError::ConnectionClosedByPeer); - } + let (fixed_header, body) = read_frame(reader, MAX_PACKET_SIZE).await?; + Packet::decode_from_body( + fixed_header.packet_type, + &fixed_header, + &mut body.as_slice(), + ) +} - let mut cursor = &header_buf[..n]; - let fixed_header = FixedHeader::decode(&mut cursor)?; +pub(crate) async fn read_frame( + reader: &mut WasmReader, + max_packet_size: u32, +) -> Result<(FixedHeader, Vec)> { + let first_byte = read_byte(reader).await?; + let packet_type_value = first_byte >> 4; + let packet_type = PacketType::from_u8(packet_type_value) + .ok_or(MqttError::InvalidPacketType(packet_type_value))?; - let remaining_length = fixed_header.remaining_length as usize; - let max_size = MAX_PACKET_SIZE as usize; - if remaining_length > max_size { - return Err(MqttError::PacketTooLarge { - size: remaining_length, - max: max_size, - }); + let mut remaining_length: u32 = 0; + let mut length_bytes = 0; + loop { + if length_bytes == MAX_REMAINING_LENGTH_BYTES { + return Err(MqttError::MalformedPacket( + "Remaining Length exceeds four bytes".to_string(), + )); + } + let byte = read_byte(reader).await?; + remaining_length |= u32::from(byte & 0x7F) << (7 * length_bytes); + length_bytes += 1; + if byte & 0x80 == 0 { + break; + } } - let mut body_buf = vec![0u8; remaining_length]; - - if remaining_length > 0 { - let bytes_read = if cursor.remaining() > 0 { - let available = cursor.remaining().min(remaining_length); - body_buf[..available].copy_from_slice(&cursor[..available]); - cursor.advance(available); - available - } else { - 0 - }; - - if bytes_read < remaining_length { - reader.read_exact(&mut body_buf[bytes_read..]).await?; - } + let header_len = 1 + length_bytes; + let total = usize::try_from(remaining_length) + .ok() + .and_then(|len| len.checked_add(header_len)) + .ok_or_else(|| MqttError::MalformedPacket("Remaining Length overflow".to_string()))?; + let max = usize::try_from(max_packet_size).unwrap_or(usize::MAX); + if total > max { + return Err(MqttError::PacketTooLarge { size: total, max }); } - let mut body = &body_buf[..]; - Packet::decode_from_body(fixed_header.packet_type, &fixed_header, &mut body) + let mut body = vec![0u8; total - header_len]; + read_exact(reader, &mut body).await?; + + let fixed_header = FixedHeader::new(packet_type, first_byte & 0x0F, remaining_length); + Ok((fixed_header, body)) } -#[allow(async_fn_in_trait)] -pub trait ReadExact { - async fn read_exact(&mut self, buf: &mut [u8]) -> Result<()>; +async fn read_byte(reader: &mut WasmReader) -> Result { + let mut byte = [0u8; 1]; + read_exact(reader, &mut byte).await?; + Ok(byte[0]) } -impl ReadExact for WasmReader { - async fn read_exact(&mut self, buf: &mut [u8]) -> Result<()> { - let mut total_read = 0; - while total_read < buf.len() { - let n = self.read(&mut buf[total_read..]).await?; - if n == 0 { - return Err(MqttError::ConnectionClosedByPeer); - } - total_read += n; +async fn read_exact(reader: &mut WasmReader, buf: &mut [u8]) -> Result<()> { + let mut total_read = 0; + while total_read < buf.len() { + let n = reader.read(&mut buf[total_read..]).await?; + if n == 0 { + return Err(MqttError::ConnectionClosedByPeer); } - Ok(()) + total_read += n; } + Ok(()) } diff --git a/crates/mqtt5-wasm/src/transport/message_port.rs b/crates/mqtt5-wasm/src/transport/message_port.rs index 134a88a1..c69e7130 100644 --- a/crates/mqtt5-wasm/src/transport/message_port.rs +++ b/crates/mqtt5-wasm/src/transport/message_port.rs @@ -28,11 +28,11 @@ pub struct MessagePortWriter { port: MessagePort, connected: Arc, msg_tx: Option>>, - _closure: Option>, + on_message: Option>, } impl MessagePortReader { - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new(rx: mpsc::UnboundedReceiver>, connected: Arc) -> Self { Self { rx, @@ -84,7 +84,7 @@ impl MessagePortReader { } impl MessagePortWriter { - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new( port: MessagePort, connected: Arc, @@ -94,13 +94,16 @@ impl MessagePortWriter { port, connected, msg_tx, - _closure: None, + on_message: None, } } /// # Errors /// Returns an error if writing to the port fails. pub fn write(&mut self, buf: &[u8]) -> Result<()> { + if !self.is_connected() { + return Err(MqttError::NotConnected); + } let array = js_sys::Uint8Array::from(buf); self.port .post_message(&array.buffer()) @@ -131,13 +134,15 @@ impl Drop for MessagePortWriter { tx.close_channel(); } self.connected.store(false, Ordering::SeqCst); - self.port.set_onmessage(None); + if self.on_message.take().is_some() { + self.port.set_onmessage(None); + } self.port.close(); } } impl MessagePortTransport { - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new(port: MessagePort) -> Self { Self { port, @@ -168,7 +173,7 @@ impl MessagePortTransport { port, connected: self.connected, msg_tx: Some(msg_tx), - _closure: Some(closure), + on_message: Some(closure), }; Ok((reader, writer)) diff --git a/crates/mqtt5-wasm/src/transport/websocket.rs b/crates/mqtt5-wasm/src/transport/websocket.rs index 99d62273..3957f933 100644 --- a/crates/mqtt5-wasm/src/transport/websocket.rs +++ b/crates/mqtt5-wasm/src/transport/websocket.rs @@ -21,10 +21,19 @@ pub struct WasmWebSocketTransport { } struct ClosureBundle { - _onmessage: Closure, - _onopen: Closure, - _onerror: Closure, - _onclose: Closure, + onmessage: Closure, + onopen: Closure, + onerror: Closure, + onclose: Closure, +} + +impl ClosureBundle { + fn attach(&self, ws: &WebSocket) { + ws.set_onmessage(Some(self.onmessage.as_ref().unchecked_ref())); + ws.set_onopen(Some(self.onopen.as_ref().unchecked_ref())); + ws.set_onerror(Some(self.onerror.as_ref().unchecked_ref())); + ws.set_onclose(Some(self.onclose.as_ref().unchecked_ref())); + } } pub struct WasmReader { @@ -37,7 +46,7 @@ pub struct WasmWriter { ws: WebSocket, connected: Arc, msg_tx: mpsc::UnboundedSender>, - _closures: ClosureBundle, + closures: Option, } impl WasmReader { @@ -77,6 +86,9 @@ impl WasmWriter { /// # Errors /// Returns an error if the WebSocket send operation fails. pub fn write(&mut self, buf: &[u8]) -> Result<()> { + if !self.is_connected() { + return Err(MqttError::NotConnected); + } self.ws .send_with_u8_array(buf) .map_err(|e| MqttError::Io(format!("WebSocket send failed: {e:?}")))?; @@ -86,14 +98,20 @@ impl WasmWriter { /// # Errors /// This method does not currently return errors but uses Result for API consistency. pub fn close(&mut self) -> Result<()> { + self.shutdown(); + Ok(()) + } + + fn shutdown(&mut self) { self.msg_tx.close_channel(); - self.ws.set_onmessage(None); - self.ws.set_onopen(None); - self.ws.set_onerror(None); - self.ws.set_onclose(None); - self.ws.close().ok(); self.connected.store(false, Ordering::SeqCst); - Ok(()) + if self.closures.take().is_some() { + self.ws.set_onmessage(None); + self.ws.set_onopen(None); + self.ws.set_onerror(None); + self.ws.set_onclose(None); + } + self.ws.close().ok(); } #[must_use] @@ -104,18 +122,12 @@ impl WasmWriter { impl Drop for WasmWriter { fn drop(&mut self) { - self.msg_tx.close_channel(); - self.connected.store(false, Ordering::SeqCst); - self.ws.set_onmessage(None); - self.ws.set_onopen(None); - self.ws.set_onerror(None); - self.ws.set_onclose(None); - self.ws.close().ok(); + self.shutdown(); } } impl WasmWebSocketTransport { - #[allow(clippy::must_use_candidate)] + #[must_use] pub fn new(url: impl Into) -> Self { Self { url: url.into(), @@ -146,7 +158,7 @@ impl WasmWebSocketTransport { ws, connected: self.connected, msg_tx, - _closures: closures, + closures: Some(closures), }; Ok((reader, writer)) @@ -165,13 +177,18 @@ impl Transport for WasmWebSocketTransport { let (result_tx, result_rx) = oneshot::channel(); let msg_tx_message = msg_tx.clone(); + let connected_message = self.connected.clone(); + let ws_message = ws.clone(); let onmessage = Closure::new(move |e: MessageEvent| { if let Ok(abuf) = e.data().dyn_into::() { let array = js_sys::Uint8Array::new(&abuf); let vec = array.to_vec(); let _ = msg_tx_message.unbounded_send(vec); } else { - web_sys::console::warn_1(&"WebSocket received non-ArrayBuffer message".into()); + tracing::warn!("WebSocket received a non-binary data frame, closing connection"); + connected_message.store(false, Ordering::SeqCst); + msg_tx_message.close_channel(); + ws_message.close().ok(); } }); @@ -206,20 +223,18 @@ impl Transport for WasmWebSocketTransport { msg_tx_close.close_channel(); }); - ws.set_onmessage(Some(onmessage.as_ref().unchecked_ref())); - ws.set_onopen(Some(onopen.as_ref().unchecked_ref())); - ws.set_onerror(Some(onerror.as_ref().unchecked_ref())); - ws.set_onclose(Some(onclose.as_ref().unchecked_ref())); + let closures = ClosureBundle { + onmessage, + onopen, + onerror, + onclose, + }; + closures.attach(&ws); self.ws = Some(ws.clone()); self.rx = Some(msg_rx); self.tx = Some(msg_tx); - self.closures = Some(ClosureBundle { - _onmessage: onmessage, - _onopen: onopen, - _onerror: onerror, - _onclose: onclose, - }); + self.closures = Some(closures); let result = result_rx .await diff --git a/crates/mqtt5-wasm/tests/conformance_client.rs b/crates/mqtt5-wasm/tests/conformance_client.rs new file mode 100644 index 00000000..127dc30d --- /dev/null +++ b/crates/mqtt5-wasm/tests/conformance_client.rs @@ -0,0 +1,1606 @@ +#![cfg(target_arch = "wasm32")] + +use bytes::BytesMut; +use mqtt5_protocol::packet::auth::AuthPacket; +use mqtt5_protocol::packet::connack::ConnAckPacket; +use mqtt5_protocol::packet::disconnect::DisconnectPacket; +use mqtt5_protocol::packet::puback::PubAckPacket; +use mqtt5_protocol::packet::pubcomp::PubCompPacket; +use mqtt5_protocol::packet::publish::PublishPacket; +use mqtt5_protocol::packet::pubrec::PubRecPacket; +use mqtt5_protocol::packet::pubrel::PubRelPacket; +use mqtt5_protocol::packet::suback::SubAckPacket; +use mqtt5_protocol::packet::{FixedHeader, MqttPacket, Packet}; +use mqtt5_protocol::protocol::v5::properties::{PropertyId, PropertyValue}; +use mqtt5_protocol::protocol::v5::reason_codes::ReasonCode; +use mqtt5_protocol::QoS; +use mqtt5_wasm::{WasmConnectOptions, WasmMqttClient, WasmPublishOptions, WasmSubscribeOptions}; +use std::cell::{Cell, RefCell}; +use std::rc::Rc; +use wasm_bindgen::prelude::*; +use wasm_bindgen::JsCast; +use wasm_bindgen_futures::{spawn_local, JsFuture}; +use wasm_bindgen_test::wasm_bindgen_test; +use web_sys::{Event, MessageChannel, MessageEvent, MessagePort}; + +async fn sleep(ms: i32) { + let promise = js_sys::Promise::new(&mut |resolve, _| { + let set_timeout = js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("setTimeout")) + .unwrap() + .unchecked_into::(); + set_timeout + .call2(&JsValue::NULL, &resolve, &JsValue::from(ms)) + .unwrap(); + }); + JsFuture::from(promise).await.unwrap(); +} + +fn encode(packet: &impl MqttPacket) -> Vec { + let mut buf = BytesMut::new(); + packet.encode(&mut buf).unwrap(); + buf.to_vec() +} + +struct Frame { + first_byte: u8, + packet: Packet, +} + +fn take_frame(inbox: &mut Vec) -> Option { + let mut cursor = &inbox[..]; + let header = FixedHeader::decode(&mut cursor).ok()?; + let header_len = inbox.len() - cursor.len(); + let total = header_len + header.remaining_length as usize; + if inbox.len() < total { + return None; + } + let first_byte = inbox[0]; + let mut body = &inbox[header_len..total]; + let packet = Packet::decode_from_body(header.packet_type, &header, &mut body).unwrap(); + inbox.drain(..total); + Some(Frame { first_byte, packet }) +} + +struct FakeBroker { + port: MessagePort, + inbox: Rc>>, + closed: Rc>, + on_message: Closure, + on_close: Closure, +} + +impl FakeBroker { + fn new() -> (Self, MessagePort) { + let channel = MessageChannel::new().unwrap(); + let port = channel.port1(); + let inbox = Rc::new(RefCell::new(Vec::new())); + let closed = Rc::new(Cell::new(false)); + let inbox_in = Rc::clone(&inbox); + let on_message = Closure::::new(move |event: MessageEvent| { + let data = js_sys::Uint8Array::new(&event.data()); + inbox_in.borrow_mut().extend(data.to_vec()); + }); + port.add_event_listener_with_callback("message", on_message.as_ref().unchecked_ref()) + .unwrap(); + let closed_in = Rc::clone(&closed); + let on_close = Closure::::new(move |_: Event| closed_in.set(true)); + port.add_event_listener_with_callback("close", on_close.as_ref().unchecked_ref()) + .unwrap(); + port.start(); + ( + Self { + port, + inbox, + closed, + on_message, + on_close, + }, + channel.port2(), + ) + } + + fn send_raw(&self, bytes: &[u8]) { + let array = js_sys::Uint8Array::from(bytes); + self.port.post_message(&array.buffer()).unwrap(); + } + + fn send(&self, packet: &impl MqttPacket) { + self.send_raw(&encode(packet)); + } + + fn take_frame(&self) -> Option { + take_frame(&mut self.inbox.borrow_mut()) + } + + async fn next_frame(&self) -> Frame { + for _ in 0..200 { + if let Some(frame) = self.take_frame() { + return frame; + } + sleep(5).await; + } + panic!("client sent no packet"); + } + + async fn next_packet(&self) -> Packet { + self.next_frame().await.packet + } + + async fn next_non_ping(&self) -> Frame { + loop { + let frame = self.next_frame().await; + if !matches!(frame.packet, Packet::PingReq) { + return frame; + } + } + } + + async fn expect_silence(&self, ms: i32) { + sleep(ms).await; + while let Some(frame) = self.take_frame() { + assert!( + matches!(frame.packet, Packet::PingReq), + "unexpected packet from client: {:?}", + frame.packet + ); + } + } + + async fn wait_closed(&self) -> bool { + for _ in 0..200 { + if self.closed.get() { + return true; + } + sleep(5).await; + } + false + } + + async fn expect_disconnect(&self, reason: ReasonCode) { + match self.next_non_ping().await.packet { + Packet::Disconnect(disconnect) => assert_eq!(disconnect.reason_code, reason), + other => panic!("expected DISCONNECT {reason:?}, got {other:?}"), + } + assert!( + self.wait_closed().await, + "client did not close the connection" + ); + } + + async fn ack_subscribe(&self) { + match self.next_non_ping().await.packet { + Packet::Subscribe(subscribe) => { + self.send( + &SubAckPacket::new(subscribe.packet_id).add_granted_qos(QoS::ExactlyOnce), + ); + } + other => panic!("expected SUBSCRIBE, got {other:?}"), + } + } +} + +impl Drop for FakeBroker { + fn drop(&mut self) { + self.port + .remove_event_listener_with_callback( + "message", + self.on_message.as_ref().unchecked_ref(), + ) + .unwrap(); + self.port + .remove_event_listener_with_callback("close", self.on_close.as_ref().unchecked_ref()) + .unwrap(); + } +} + +type Outcome = Rc>>>; + +fn spawn_outcome(future: F) -> Outcome +where + F: std::future::Future> + 'static, +{ + let outcome: Outcome = Rc::new(RefCell::new(None)); + let slot = Rc::clone(&outcome); + spawn_local(async move { + let result = future.await; + *slot.borrow_mut() = Some(result); + }); + outcome +} + +async fn settle(outcome: &Outcome) -> Result { + for _ in 0..400 { + if let Some(result) = outcome.borrow_mut().take() { + return result; + } + sleep(5).await; + } + panic!("operation did not settle"); +} + +fn is_pending(outcome: &Outcome) -> bool { + outcome.borrow().is_none() +} + +async fn assert_rejected_silently(outcome: &Outcome, broker: &FakeBroker) { + broker.expect_silence(60).await; + let result = settle(outcome).await; + assert!( + result.is_err(), + "operation should be rejected, got {result:?}" + ); +} + +fn noop() -> js_sys::Function { + js_sys::Function::new_no_args("") +} + +fn recorder() -> (js_sys::Function, Rc>>) { + let topics = Rc::new(RefCell::new(Vec::new())); + let sink = Rc::clone(&topics); + let callback = Closure::::new( + move |topic: JsValue, _payload: JsValue, _props: JsValue| { + sink.borrow_mut() + .push(topic.as_string().unwrap_or_default()); + }, + ); + (callback.into_js_value().unchecked_into(), topics) +} + +fn value_recorder() -> (js_sys::Function, Rc>>) { + let values = Rc::new(RefCell::new(Vec::new())); + let sink = Rc::clone(&values); + let callback = Closure::::new(move |value: JsValue| { + sink.borrow_mut().push(value); + }); + (callback.into_js_value().unchecked_into(), values) +} + +fn success() -> ConnAckPacket { + ConnAckPacket::new(false, ReasonCode::Success) +} + +async fn open_session( + client: &Rc, + options: WasmConnectOptions, + connack: ConnAckPacket, +) -> (FakeBroker, Result<(), JsValue>, Packet) { + let (broker, client_port) = FakeBroker::new(); + let connecting = { + let client = Rc::clone(client); + spawn_outcome(async move { + client + .connect_message_port_with_options(client_port, &options) + .await + }) + }; + let connect = broker.next_packet().await; + assert!( + matches!(connect, Packet::Connect(_)), + "expected CONNECT, got {connect:?}" + ); + broker.send(&connack); + let result = settle(&connecting).await; + (broker, result, connect) +} + +async fn connect_with( + options: WasmConnectOptions, + connack: ConnAckPacket, +) -> (Rc, FakeBroker) { + let client = Rc::new(WasmMqttClient::new("wasm-conformance".to_string())); + let (broker, result, _) = open_session(&client, options, connack).await; + result.expect("connect failed"); + (client, broker) +} + +async fn connect_default() -> (Rc, FakeBroker) { + connect_with(WasmConnectOptions::new(), success()).await +} + +fn publish_options(qos: u8) -> WasmPublishOptions { + let mut options = WasmPublishOptions::new(); + options.set_qos(qos); + options +} + +fn publish_with( + client: &Rc, + topic: &str, + payload: &[u8], + options: WasmPublishOptions, +) -> Outcome<()> { + let client = Rc::clone(client); + let topic = topic.to_string(); + let payload = payload.to_vec(); + spawn_outcome(async move { + client + .publish_with_options(&topic, &payload, &options) + .await + }) +} + +fn subscribe_with( + client: &Rc, + filter: &str, + options: WasmSubscribeOptions, +) -> Outcome { + let client = Rc::clone(client); + let filter = filter.to_string(); + spawn_outcome(async move { + client + .subscribe_with_options(&filter, noop(), &options) + .await + }) +} + +fn inbound_publish(topic: &str, qos: QoS, packet_id: Option) -> PublishPacket { + let packet = PublishPacket::new(topic.to_string(), b"payload".to_vec(), qos); + match packet_id { + Some(id) => packet.with_packet_id(id), + None => packet, + } +} + +async fn subscribed_client( + options: WasmConnectOptions, +) -> (Rc, FakeBroker, Rc>>) { + let (client, broker) = connect_with(options, success()).await; + let (callback, topics) = recorder(); + client.subscribe_with_callback("#", callback).await.unwrap(); + broker.ack_subscribe().await; + (client, broker, topics) +} + +async fn wait_for_count(items: &Rc>>, count: usize) { + for _ in 0..200 { + if items.borrow().len() >= count { + return; + } + sleep(5).await; + } +} + +#[wasm_bindgen_test] +async fn mqtt_4_7_0_1_publish_topic_with_wildcard_rejected() { + let (client, broker) = connect_default().await; + assert!(client.publish("a/+", b"x").await.is_err()); + assert!(client.publish("a/#", b"x").await.is_err()); + let outcome = publish_with(&client, "a/+/b", b"x", publish_options(1)); + assert_rejected_silently(&outcome, &broker).await; + assert!(client.publish_qos1("#", b"x", noop()).await.is_err()); + assert!(client.publish_qos2("+", b"x", noop()).await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn mqtt_4_7_3_1_empty_topic_without_alias_rejected() { + let (client, broker) = connect_default().await; + assert!(client.publish("", b"x").await.is_err()); + let outcome = publish_with(&client, "", b"x", publish_options(1)); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn mqtt_4_7_3_1_empty_topic_with_unmapped_alias_rejected() { + let (client, broker) = connect_with( + WasmConnectOptions::new(), + success().with_topic_alias_maximum(5), + ) + .await; + let mut options = publish_options(0); + options.set_topic_alias(Some(2)); + let outcome = publish_with(&client, "", b"x", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn mqtt_4_7_3_2_publish_topic_with_null_rejected() { + let (client, broker) = connect_default().await; + assert!(client.publish("a\0b", b"x").await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_2_14_wildcard_response_topic_rejected() { + let (client, broker) = connect_default().await; + let mut options = publish_options(0); + options.set_response_topic(Some("reply/#".to_string())); + let outcome = publish_with(&client, "a", b"x", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_2_8_topic_alias_zero_rejected() { + let (client, broker) = connect_with( + WasmConnectOptions::new(), + success().with_topic_alias_maximum(5), + ) + .await; + let mut options = publish_options(0); + options.set_topic_alias(Some(0)); + let outcome = publish_with(&client, "a", b"x", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_18_topic_alias_without_server_maximum_rejected() { + let (client, broker) = connect_default().await; + let mut options = publish_options(0); + options.set_topic_alias(Some(1)); + let outcome = publish_with(&client, "a", b"x", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_2_9_topic_alias_above_server_maximum_rejected() { + let (client, broker) = connect_with( + WasmConnectOptions::new(), + success().with_topic_alias_maximum(2), + ) + .await; + let mut above = publish_options(0); + above.set_topic_alias(Some(3)); + let outcome = publish_with(&client, "a", b"x", above); + assert_rejected_silently(&outcome, &broker).await; + + let mut within = publish_options(0); + within.set_topic_alias(Some(2)); + settle(&publish_with(&client, "a", b"x", within)) + .await + .unwrap(); + match broker.next_non_ping().await.packet { + Packet::Publish(publish) => assert_eq!(publish.topic_alias(), Some(2)), + other => panic!("expected PUBLISH, got {other:?}"), + } + let mut reuse = publish_options(0); + reuse.set_topic_alias(Some(2)); + settle(&publish_with(&client, "", b"y", reuse)) + .await + .unwrap(); + match broker.next_non_ping().await.packet { + Packet::Publish(publish) => { + assert_eq!(publish.topic_name, ""); + assert_eq!(publish.topic_alias(), Some(2)); + } + other => panic!("expected PUBLISH, got {other:?}"), + } +} + +#[wasm_bindgen_test] +async fn mqtt_4_7_1_1_invalid_multi_level_wildcard_filter_rejected() { + let (client, broker) = connect_default().await; + assert!(client.subscribe("a/#/b").await.is_err()); + assert!(client.subscribe_with_callback("a#", noop()).await.is_err()); + let outcome = subscribe_with(&client, "sport/#/x", WasmSubscribeOptions::new()); + assert_rejected_silently(&outcome, &broker).await; + assert!(client.unsubscribe("a/#/b").await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn mqtt_4_7_1_2_partial_level_single_wildcard_rejected() { + let (client, broker) = connect_default().await; + assert!(client.subscribe("a+/b").await.is_err()); + assert!(client.unsubscribe("a/b+").await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn mqtt_4_7_3_1_empty_filter_rejected() { + let (client, broker) = connect_default().await; + assert!(client.subscribe("").await.is_err()); + assert!(client.unsubscribe("").await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn mqtt_4_8_2_1_shared_subscription_without_share_name_rejected() { + let (client, broker) = connect_default().await; + assert!(client.subscribe("$share//a").await.is_err()); + assert!(client.subscribe("$share/group").await.is_err()); + assert!(client.subscribe("$share/group/").await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn mqtt_4_8_2_2_share_name_with_wildcard_rejected() { + let (client, broker) = connect_default().await; + assert!(client.subscribe("$share/g+/a").await.is_err()); + assert!(client.subscribe("$share/#/a").await.is_err()); + assert!(client.unsubscribe("$share/g#/a").await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_8_3_4_no_local_on_shared_subscription_rejected() { + let (client, broker) = connect_default().await; + let mut options = WasmSubscribeOptions::new(); + options.set_no_local(true); + let outcome = subscribe_with(&client, "$share/g/a", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn section_3_8_2_1_2_subscription_identifier_zero_rejected() { + let (client, broker) = connect_default().await; + let mut options = WasmSubscribeOptions::new(); + options.set_subscription_identifier(Some(0)); + let outcome = subscribe_with(&client, "a", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn section_3_2_2_3_11_wildcard_subscription_unavailable_honoured() { + let mut connack = success(); + connack + .properties + .add( + PropertyId::WildcardSubscriptionAvailable, + PropertyValue::Byte(0), + ) + .unwrap(); + let (client, broker) = connect_with(WasmConnectOptions::new(), connack).await; + assert!(client.subscribe("a/+").await.is_err()); + assert!(client.subscribe_with_callback("a/#", noop()).await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn section_3_2_2_3_13_shared_subscription_unavailable_honoured() { + let mut connack = success(); + connack + .properties + .add( + PropertyId::SharedSubscriptionAvailable, + PropertyValue::Byte(0), + ) + .unwrap(); + let (client, broker) = connect_with(WasmConnectOptions::new(), connack).await; + assert!(client.subscribe("$share/g/a").await.is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn section_3_2_2_3_12_subscription_identifier_unavailable_honoured() { + let mut connack = success(); + connack + .properties + .add( + PropertyId::SubscriptionIdentifierAvailable, + PropertyValue::Byte(0), + ) + .unwrap(); + let (client, broker) = connect_with(WasmConnectOptions::new(), connack).await; + let mut options = WasmSubscribeOptions::new(); + options.set_subscription_identifier(Some(7)); + let outcome = subscribe_with(&client, "a", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_14_retain_unavailable_honoured() { + let (client, broker) = connect_with( + WasmConnectOptions::new(), + success().with_retain_available(false), + ) + .await; + let mut options = publish_options(0); + options.set_retain(true); + let outcome = publish_with(&client, "a", b"x", options); + assert_rejected_silently(&outcome, &broker).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_14_capabilities_refreshed_on_reconnect() { + let client = Rc::new(WasmMqttClient::new("refresh".to_string())); + let (broker, result, _) = open_session( + &client, + WasmConnectOptions::new(), + success().with_retain_available(false), + ) + .await; + result.unwrap(); + let mut retained = publish_options(0); + retained.set_retain(true); + let outcome = publish_with(&client, "a", b"x", retained); + assert_rejected_silently(&outcome, &broker).await; + client.disconnect().await.unwrap(); + + let (broker, result, _) = open_session(&client, WasmConnectOptions::new(), success()).await; + result.unwrap(); + let mut retained = publish_options(0); + retained.set_retain(true); + settle(&publish_with(&client, "a", b"x", retained)) + .await + .unwrap(); + match broker.next_non_ping().await.packet { + Packet::Publish(publish) => assert!(publish.retain), + other => panic!("expected PUBLISH, got {other:?}"), + } +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_11_maximum_qos_honoured() { + let (client, broker) = + connect_with(WasmConnectOptions::new(), success().with_maximum_qos(0)).await; + let outcome = publish_with(&client, "a", b"x", publish_options(1)); + let qos1 = client.publish_qos1("a", b"x", noop()).await; + let qos2 = client.publish_qos2("a", b"x", noop()).await; + sleep(60).await; + while let Some(frame) = broker.take_frame() { + if let Packet::Publish(publish) = frame.packet { + assert_eq!(publish.qos, QoS::AtMostOnce, "PUBLISH exceeds Maximum QoS"); + } + } + assert!(settle(&outcome).await.is_err()); + assert!(qos1.is_err()); + assert!(qos2.is_err()); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_15_server_maximum_packet_size_honoured() { + let (client, broker) = connect_with( + WasmConnectOptions::new(), + success().with_maximum_packet_size(64), + ) + .await; + let big = vec![0u8; 200]; + assert!(client.publish("a", &big).await.is_err()); + let outcome = publish_with(&client, "a", &big, publish_options(1)); + assert_rejected_silently(&outcome, &broker).await; + let long_filter = "f".repeat(100); + assert!(client.subscribe(&long_filter).await.is_err()); + assert!(client.unsubscribe(&long_filter).await.is_err()); + broker.expect_silence(40).await; + client.publish("a", b"small").await.unwrap(); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Publish(_) + )); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_21_server_keep_alive_used() { + let mut options = WasmConnectOptions::new(); + options.set_keep_alive(60); + let (client, broker) = connect_with(options, success().with_server_keep_alive(1)).await; + let mut pinged = false; + for _ in 0..30 { + sleep(50).await; + if let Some(frame) = broker.take_frame() { + pinged = matches!(frame.packet, Packet::PingReq); + break; + } + } + client.disconnect().await.unwrap(); + assert!(pinged, "client ignored the Server Keep Alive"); +} + +#[wasm_bindgen_test] +async fn section_3_1_2_10_keep_alive_zero_sends_no_pingreq() { + let mut options = WasmConnectOptions::new(); + options.set_keep_alive(0); + let (client, broker) = connect_with(options, success()).await; + sleep(100).await; + assert!( + broker.take_frame().is_none(), + "keep alive 0 must not trigger PINGREQ" + ); + assert!(client.is_connected()); + client.disconnect().await.unwrap(); +} + +#[wasm_bindgen_test] +async fn mqtt_3_1_3_2_assigned_client_identifier_used_for_session() { + let client = Rc::new(WasmMqttClient::new(String::new())); + let mut connack = success(); + connack + .properties + .add( + PropertyId::AssignedClientIdentifier, + PropertyValue::Utf8String("server-assigned".to_string()), + ) + .unwrap(); + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(300)); + let (broker, result, _) = open_session(&client, options, connack).await; + result.unwrap(); + client.disconnect().await.unwrap(); + assert!(matches!(broker.next_packet().await, Packet::Disconnect(_))); + + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(300)); + let (_broker, result, connect) = open_session( + &client, + options, + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + result.unwrap(); + match connect { + Packet::Connect(connect) => assert_eq!(connect.client_id, "server-assigned"), + other => panic!("expected CONNECT, got {other:?}"), + } +} + +#[wasm_bindgen_test] +async fn section_3_1_2_11_connect_carries_requested_properties() { + let client = Rc::new(WasmMqttClient::new("props".to_string())); + let mut options = WasmConnectOptions::new(); + options.set_request_problem_information(Some(false)); + options.set_request_response_information(Some(true)); + options.add_user_property("k".to_string(), "v".to_string()); + let (_broker, result, connect) = open_session(&client, options, success()).await; + result.unwrap(); + match connect { + Packet::Connect(connect) => { + assert_eq!( + connect.properties.get_request_problem_information(), + Some(false) + ); + assert_eq!( + connect.properties.get_request_response_information(), + Some(true) + ); + assert!(connect.properties.contains(PropertyId::UserProperty)); + } + other => panic!("expected CONNECT, got {other:?}"), + } +} + +#[wasm_bindgen_test] +async fn mqtt_4_12_0_7_no_auth_without_authentication_method() { + let client = Rc::new(WasmMqttClient::new("auth".to_string())); + let mut with_method = WasmConnectOptions::new(); + with_method.set_authentication_method(Some("SCRAM".to_string())); + let (broker, result, _) = open_session( + &client, + with_method, + success().with_authentication_method("SCRAM".to_string()), + ) + .await; + result.unwrap(); + client.disconnect().await.unwrap(); + assert!(matches!(broker.next_packet().await, Packet::Disconnect(_))); + + let (broker, result, _) = open_session(&client, WasmConnectOptions::new(), success()).await; + result.unwrap(); + assert!(client.respond_auth(b"data").is_err()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn section_4_12_unexpected_auth_is_protocol_error() { + let (_client, broker) = connect_default().await; + let mut auth = AuthPacket::new(ReasonCode::ContinueAuthentication); + auth.properties + .set_authentication_method("SCRAM".to_string()); + broker.send(&auth); + broker.expect_disconnect(ReasonCode::ProtocolError).await; +} + +#[wasm_bindgen_test] +async fn mqtt_2_2_1_3_packet_identifier_not_reused_while_in_flight() { + let (client, broker) = connect_default().await; + let held = client.publish_qos1("held", b"x", noop()).await.unwrap(); + assert!(matches!(broker.next_packet().await, Packet::Publish(_))); + let mut issued = 0u32; + while issued < 65_600 { + let mut ids = Vec::with_capacity(2000); + for _ in 0..2000 { + let id = client.publish_qos1("t", b"", noop()).await.unwrap(); + assert_ne!(id, held, "packet identifier {held} reused while in flight"); + ids.push(id); + } + issued += 2000; + let mut acks = Vec::with_capacity(ids.len() * 4); + for _ in &ids { + let frame = broker.next_frame().await; + match frame.packet { + Packet::Publish(publish) => { + acks.extend(encode(&PubAckPacket::new(publish.packet_id.unwrap()))); + } + other => panic!("expected PUBLISH, got {other:?}"), + } + } + for chunk in acks.chunks(4) { + broker.send_raw(chunk); + } + sleep(20).await; + } +} + +#[wasm_bindgen_test] +async fn mqtt_4_4_0_1_unacknowledged_messages_resent_on_session_resume() { + let client = Rc::new(WasmMqttClient::new("resume".to_string())); + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(3600)); + let (broker, result, _) = open_session(&client, options, success()).await; + result.unwrap(); + + let first = publish_with(&client, "a", b"1", publish_options(1)); + let second = publish_with(&client, "b", b"2", publish_options(2)); + let third = publish_with(&client, "c", b"3", publish_options(1)); + let mut ids = Vec::new(); + for _ in 0..3 { + match broker.next_non_ping().await.packet { + Packet::Publish(publish) => ids.push(publish.packet_id.unwrap()), + other => panic!("expected PUBLISH, got {other:?}"), + } + } + broker.send(&PubRecPacket::new(ids[1])); + match broker.next_non_ping().await.packet { + Packet::PubRel(pubrel) => assert_eq!(pubrel.packet_id, ids[1]), + other => panic!("expected PUBREL, got {other:?}"), + } + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + sleep(50).await; + assert!(!client.is_connected()); + + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(3600)); + let (broker, result, _) = open_session( + &client, + options, + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + result.unwrap(); + + let resent: Vec = { + let mut frames = Vec::new(); + for _ in 0..3 { + frames.push(broker.next_non_ping().await); + } + frames + }; + match &resent[0].packet { + Packet::Publish(publish) => { + assert_eq!(publish.packet_id, Some(ids[0])); + assert!(publish.dup, "resent PUBLISH must have DUP=1"); + assert_eq!(publish.topic_name, "a"); + } + other => panic!("expected resent PUBLISH a, got {other:?}"), + } + match &resent[1].packet { + Packet::PubRel(pubrel) => assert_eq!(pubrel.packet_id, ids[1]), + other => panic!("expected resent PUBREL, got {other:?}"), + } + match &resent[2].packet { + Packet::Publish(publish) => { + assert_eq!(publish.packet_id, Some(ids[2])); + assert_eq!(resent[2].first_byte & 0x08, 0x08); + assert_eq!(publish.topic_name, "c"); + } + other => panic!("expected resent PUBLISH c, got {other:?}"), + } + broker.send(&PubAckPacket::new(ids[0])); + broker.send(&PubCompPacket::new(ids[1])); + broker.send(&PubAckPacket::new(ids[2])); + settle(&first).await.unwrap(); + settle(&second).await.unwrap(); + settle(&third).await.unwrap(); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_4_session_present_without_session_state_closes() { + let client = Rc::new(WasmMqttClient::new("fresh".to_string())); + let (broker, result, _) = open_session( + &client, + WasmConnectOptions::new(), + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + assert!( + result.is_err(), + "Session Present=1 without session state must fail" + ); + assert!( + broker.wait_closed().await, + "client did not close the connection" + ); + assert!(!client.is_connected()); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_4_fresh_client_clean_start_zero_session_present_closes() { + let client = Rc::new(WasmMqttClient::new("fresh-resume".to_string())); + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(3600)); + let (broker, result, _) = open_session( + &client, + options, + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + assert!( + result.is_err(), + "Session Present=1 without session state must fail" + ); + assert!( + broker.wait_closed().await, + "client did not close the connection" + ); + assert!(!client.is_connected()); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_4_resume_existing_session_opt_in_accepts_session_present() { + let client = Rc::new(WasmMqttClient::new("opt-in-resume".to_string())); + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(3600)); + options.set_resume_existing_session(true); + let (broker, result, _) = open_session( + &client, + options, + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + result.expect("opt-in resume must accept Session Present=1"); + assert!(client.is_connected()); + broker.expect_silence(60).await; + assert!(!broker.closed.get()); + + broker.send(&inbound_publish("held/topic", QoS::AtLeastOnce, Some(11))); + match broker.next_non_ping().await.packet { + Packet::PubAck(puback) => assert_eq!(puback.packet_id, 11), + other => panic!("expected PUBACK, got {other:?}"), + } +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_4_resume_existing_session_still_rejects_clean_start() { + let client = Rc::new(WasmMqttClient::new("opt-in-clean".to_string())); + let mut options = WasmConnectOptions::new(); + options.set_resume_existing_session(true); + let (broker, result, _) = open_session( + &client, + options, + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + assert!( + result.is_err(), + "Session Present=1 after Clean Start=1 must fail" + ); + assert!( + broker.wait_closed().await, + "client did not close the connection" + ); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_5_session_state_discarded_on_session_present_zero() { + let client = Rc::new(WasmMqttClient::new("discard".to_string())); + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(3600)); + let (broker, result, _) = open_session(&client, options, success()).await; + result.unwrap(); + let pending = publish_with(&client, "a", b"1", publish_options(1)); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Publish(_) + )); + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + sleep(50).await; + + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(3600)); + let (broker, result, _) = open_session(&client, options, success()).await; + result.unwrap(); + broker.expect_silence(60).await; + assert!(settle(&pending).await.is_err()); +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_4_7_server_receive_maximum_enforced() { + let (client, broker) = + connect_with(WasmConnectOptions::new(), success().with_receive_maximum(1)).await; + let first = publish_with(&client, "a", b"1", publish_options(1)); + let second = publish_with(&client, "b", b"2", publish_options(1)); + let first_id = match broker.next_non_ping().await.packet { + Packet::Publish(publish) => publish.packet_id.unwrap(), + other => panic!("expected PUBLISH, got {other:?}"), + }; + broker.expect_silence(60).await; + client.publish("qos0", b"allowed").await.unwrap(); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Publish(_) + )); + assert!(is_pending(&second)); + broker.send(&PubAckPacket::new(first_id)); + settle(&first).await.unwrap(); + let second_id = match broker.next_non_ping().await.packet { + Packet::Publish(publish) => publish.packet_id.unwrap(), + other => panic!("expected second PUBLISH, got {other:?}"), + }; + broker.send(&PubAckPacket::new(second_id)); + settle(&second).await.unwrap(); +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_4_8_non_publish_packets_not_delayed_by_quota() { + let (client, broker) = + connect_with(WasmConnectOptions::new(), success().with_receive_maximum(1)).await; + let first = publish_with(&client, "a", b"1", publish_options(1)); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Publish(_) + )); + let second = publish_with(&client, "b", b"2", publish_options(1)); + client.subscribe("x").await.unwrap(); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Subscribe(_) + )); + assert!(is_pending(&first)); + assert!(is_pending(&second)); +} + +#[wasm_bindgen_test] +async fn mqtt_4_9_0_1_send_quota_reset_on_new_connection() { + let client = Rc::new(WasmMqttClient::new("quota".to_string())); + let (broker, result, _) = open_session( + &client, + WasmConnectOptions::new(), + success().with_receive_maximum(1), + ) + .await; + result.unwrap(); + let stuck = publish_with(&client, "a", b"1", publish_options(1)); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Publish(_) + )); + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + sleep(50).await; + assert!(settle(&stuck).await.is_err()); + + let (broker, result, _) = open_session( + &client, + WasmConnectOptions::new(), + success().with_receive_maximum(1), + ) + .await; + result.unwrap(); + let fresh = publish_with(&client, "b", b"2", publish_options(1)); + let id = match broker.next_non_ping().await.packet { + Packet::Publish(publish) => publish.packet_id.unwrap(), + other => panic!("expected PUBLISH, got {other:?}"), + }; + broker.send(&PubAckPacket::new(id)); + settle(&fresh).await.unwrap(); +} + +#[wasm_bindgen_test] +async fn mqtt_4_3_2_4_puback_sent_for_inbound_qos1() { + let (_client, broker, topics) = subscribed_client(WasmConnectOptions::new()).await; + broker.send(&inbound_publish("a/b", QoS::AtLeastOnce, Some(7))); + match broker.next_non_ping().await.packet { + Packet::PubAck(puback) => assert_eq!(puback.packet_id, 7), + other => panic!("expected PUBACK, got {other:?}"), + } + wait_for_count(&topics, 1).await; + assert_eq!(topics.borrow().as_slice(), ["a/b"]); +} + +#[wasm_bindgen_test] +async fn mqtt_4_6_0_2_pubacks_sent_in_receive_order() { + let (_client, broker, _topics) = subscribed_client(WasmConnectOptions::new()).await; + for id in [30, 10, 20] { + broker.send(&inbound_publish("a", QoS::AtLeastOnce, Some(id))); + } + let mut acked = Vec::new(); + for _ in 0..3 { + match broker.next_non_ping().await.packet { + Packet::PubAck(puback) => acked.push(puback.packet_id), + other => panic!("expected PUBACK, got {other:?}"), + } + } + assert_eq!(acked, [30, 10, 20]); +} + +#[wasm_bindgen_test] +async fn mqtt_4_6_0_3_pubrecs_sent_in_receive_order() { + let (_client, broker, _topics) = subscribed_client(WasmConnectOptions::new()).await; + for id in [9, 3, 6] { + broker.send(&inbound_publish("a", QoS::ExactlyOnce, Some(id))); + } + let mut recorded = Vec::new(); + for _ in 0..3 { + match broker.next_non_ping().await.packet { + Packet::PubRec(pubrec) => recorded.push(pubrec.packet_id), + other => panic!("expected PUBREC, got {other:?}"), + } + } + assert_eq!(recorded, [9, 3, 6]); +} + +#[wasm_bindgen_test] +async fn mqtt_4_3_3_11_pubcomp_for_pubrel_and_identifier_reusable() { + let (_client, broker, topics) = subscribed_client(WasmConnectOptions::new()).await; + broker.send(&inbound_publish("first", QoS::ExactlyOnce, Some(5))); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::PubRec(_) + )); + broker.send(&inbound_publish("first", QoS::ExactlyOnce, Some(5)).with_dup(true)); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::PubRec(_) + )); + broker.send(&PubRelPacket::new(5)); + match broker.next_non_ping().await.packet { + Packet::PubComp(pubcomp) => assert_eq!(pubcomp.packet_id, 5), + other => panic!("expected PUBCOMP, got {other:?}"), + } + broker.send(&inbound_publish("second", QoS::ExactlyOnce, Some(5))); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::PubRec(_) + )); + wait_for_count(&topics, 2).await; + assert_eq!(topics.borrow().as_slice(), ["first", "second"]); +} + +#[wasm_bindgen_test] +async fn mqtt_6_0_0_2_multiple_packets_in_one_frame() { + let (_client, broker, topics) = subscribed_client(WasmConnectOptions::new()).await; + let mut frame = vec![0xD0, 0x00]; + frame.extend(encode(&inbound_publish("one", QoS::AtMostOnce, None))); + frame.extend(encode(&inbound_publish("two", QoS::AtMostOnce, None))); + broker.send_raw(&frame); + wait_for_count(&topics, 2).await; + assert_eq!(topics.borrow().as_slice(), ["one", "two"]); +} + +#[wasm_bindgen_test] +async fn mqtt_6_0_0_2_packet_split_across_frames() { + let (_client, broker, topics) = subscribed_client(WasmConnectOptions::new()).await; + let bytes = encode(&inbound_publish("split/topic", QoS::AtMostOnce, None)); + broker.send_raw(&bytes[..1]); + sleep(10).await; + broker.send_raw(&bytes[1..3]); + sleep(10).await; + broker.send_raw(&bytes[3..]); + wait_for_count(&topics, 1).await; + assert_eq!(topics.borrow().as_slice(), ["split/topic"]); +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_2_10_inbound_topic_alias_resolved() { + let mut options = WasmConnectOptions::new(); + options.set_topic_alias_maximum(Some(5)); + let (_client, broker, topics) = subscribed_client(options).await; + broker.send(&inbound_publish("alias/topic", QoS::AtMostOnce, None).with_topic_alias(5)); + broker.send(&inbound_publish("", QoS::AtMostOnce, None).with_topic_alias(5)); + wait_for_count(&topics, 2).await; + assert_eq!(topics.borrow().as_slice(), ["alias/topic", "alias/topic"]); +} + +#[wasm_bindgen_test] +async fn section_3_3_2_3_4_inbound_topic_alias_above_maximum_disconnects() { + let mut options = WasmConnectOptions::new(); + options.set_topic_alias_maximum(Some(2)); + let (_client, broker, _topics) = subscribed_client(options).await; + broker.send(&inbound_publish("a", QoS::AtMostOnce, None).with_topic_alias(3)); + broker + .expect_disconnect(ReasonCode::TopicAliasInvalid) + .await; +} + +#[wasm_bindgen_test] +async fn section_3_3_2_3_4_inbound_topic_alias_zero_disconnects() { + let mut options = WasmConnectOptions::new(); + options.set_topic_alias_maximum(Some(2)); + let (_client, broker, _topics) = subscribed_client(options).await; + broker.send(&inbound_publish("a", QoS::AtMostOnce, None).with_topic_alias(0)); + broker + .expect_disconnect(ReasonCode::TopicAliasInvalid) + .await; +} + +#[wasm_bindgen_test] +async fn section_3_3_2_3_4_inbound_unknown_alias_is_protocol_error() { + let mut options = WasmConnectOptions::new(); + options.set_topic_alias_maximum(Some(2)); + let (_client, broker, _topics) = subscribed_client(options).await; + broker.send(&inbound_publish("", QoS::AtMostOnce, None).with_topic_alias(1)); + broker.expect_disconnect(ReasonCode::ProtocolError).await; +} + +#[wasm_bindgen_test] +async fn section_3_3_4_inbound_receive_maximum_exceeded_disconnects() { + let mut options = WasmConnectOptions::new(); + options.set_receive_maximum(Some(1)); + let (_client, broker, _topics) = subscribed_client(options).await; + broker.send(&inbound_publish("a", QoS::ExactlyOnce, Some(1))); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::PubRec(_) + )); + broker.send(&inbound_publish("a", QoS::ExactlyOnce, Some(2))); + broker + .expect_disconnect(ReasonCode::ReceiveMaximumExceeded) + .await; +} + +#[wasm_bindgen_test] +async fn section_3_1_2_11_4_inbound_packet_above_maximum_size_disconnects() { + let mut options = WasmConnectOptions::new(); + options.set_maximum_packet_size(Some(128)); + let (_client, broker, topics) = subscribed_client(options).await; + let mut publish = inbound_publish("big", QoS::AtMostOnce, None); + publish.payload = vec![7u8; 400].into(); + broker.send(&publish); + broker.expect_disconnect(ReasonCode::PacketTooLarge).await; + assert!(topics.borrow().is_empty()); +} + +#[wasm_bindgen_test] +async fn mqtt_3_6_1_1_pubrel_reserved_flags_malformed() { + let (_client, broker, _topics) = subscribed_client(WasmConnectOptions::new()).await; + broker.send(&inbound_publish("a", QoS::ExactlyOnce, Some(4))); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::PubRec(_) + )); + broker.send_raw(&[0x60, 0x02, 0x00, 0x04]); + broker.expect_disconnect(ReasonCode::MalformedPacket).await; +} + +#[wasm_bindgen_test] +async fn mqtt_2_1_3_1_suback_reserved_flags_malformed() { + let (client, broker) = connect_default().await; + let id = client.subscribe("a").await.unwrap(); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Subscribe(_) + )); + let [id_high, id_low] = id.to_be_bytes(); + let bytes = [0x92, 0x04, id_high, id_low, 0x00, 0x00]; + broker.send_raw(&bytes); + broker.expect_disconnect(ReasonCode::MalformedPacket).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_14_1_1_disconnect_reserved_flags_malformed() { + let (_client, broker) = connect_default().await; + broker.send_raw(&[0xE1, 0x00]); + broker.expect_disconnect(ReasonCode::MalformedPacket).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_15_1_1_auth_reserved_flags_malformed() { + let (_client, broker) = connect_default().await; + broker.send_raw(&[0xF1, 0x00]); + broker.expect_disconnect(ReasonCode::MalformedPacket).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_1_4_publish_qos_three_malformed() { + let (_client, broker) = connect_default().await; + broker.send_raw(&[0x36, 0x06, 0x00, 0x01, b'a', 0x00, 0x01, 0x00]); + broker.expect_disconnect(ReasonCode::MalformedPacket).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_1_connack_reserved_flags_rejected() { + let client = Rc::new(WasmMqttClient::new("flags".to_string())); + let (broker, client_port) = FakeBroker::new(); + let options = WasmConnectOptions::new(); + let connecting = { + let client = Rc::clone(&client); + spawn_outcome(async move { + client + .connect_message_port_with_options(client_port, &options) + .await + }) + }; + assert!(matches!(broker.next_packet().await, Packet::Connect(_))); + broker.send_raw(&[0x20, 0x03, 0x02, 0x00, 0x00]); + assert!(settle(&connecting).await.is_err()); + assert!( + broker.wait_closed().await, + "client did not close the connection" + ); +} + +#[wasm_bindgen_test] +async fn section_3_14_server_disconnect_closes_connection() { + let (client, broker) = connect_default().await; + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + assert!( + broker.wait_closed().await, + "client kept the connection open" + ); + assert!(!client.is_connected()); + broker.expect_silence(40).await; +} + +#[wasm_bindgen_test] +async fn section_3_1_2_11_7_problem_information_zero_enforced() { + let mut options = WasmConnectOptions::new(); + options.set_request_problem_information(Some(false)); + let (client, broker) = connect_with(options, success()).await; + let id = client.subscribe("a").await.unwrap(); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Subscribe(_) + )); + broker.send( + &SubAckPacket::new(id) + .add_granted_qos(QoS::AtMostOnce) + .with_reason_string("not allowed".to_string()), + ); + broker.expect_disconnect(ReasonCode::ProtocolError).await; +} + +#[wasm_bindgen_test] +async fn mqtt_3_14_4_1_nothing_sent_after_disconnect() { + let (client, broker) = connect_with( + WasmConnectOptions::new(), + success().with_server_keep_alive(1), + ) + .await; + client.disconnect().await.unwrap(); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Disconnect(_) + )); + assert!(client.publish("a", b"x").await.is_err()); + assert!(client.subscribe("a").await.is_err()); + assert!(client.respond_auth(b"x").is_err()); + sleep(700).await; + assert!( + broker.take_frame().is_none(), + "client sent a packet after DISCONNECT" + ); +} + +#[wasm_bindgen_test] +async fn mqtt_4_3_2_2_first_qos1_send_has_dup_zero() { + let (client, broker) = connect_default().await; + let outcome = publish_with(&client, "a", b"x", publish_options(1)); + let frame = broker.next_non_ping().await; + assert_eq!( + frame.first_byte & 0x08, + 0, + "first transmission must have DUP=0" + ); + if let Packet::Publish(publish) = frame.packet { + broker.send(&PubAckPacket::new(publish.packet_id.unwrap())); + } + settle(&outcome).await.unwrap(); +} + +#[wasm_bindgen_test] +async fn section_3_4_puback_error_reason_completes_flow_without_retransmit() { + let (client, broker) = connect_default().await; + let (callback, results) = value_recorder(); + let id = client.publish_qos1("a", b"x", callback).await.unwrap(); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::Publish(_) + )); + broker.send(&PubAckPacket::new_with_reason( + id, + ReasonCode::NotAuthorized, + )); + wait_for_count(&results, 1).await; + assert_eq!(results.borrow()[0].as_f64(), Some(f64::from(0x87_u8))); + broker.expect_silence(60).await; +} + +#[wasm_bindgen(inline_js = r#" +const net = process.getBuiltinModule('node:net'); +const crypto = process.getBuiltinModule('node:crypto'); +export class WsFake { + constructor() { this.protocols = ''; this.opcodes = []; this.data = []; this.closed = false; this.sock = null; this.buf = Buffer.alloc(0); } + start() { + return new Promise((resolve) => { + this.server = net.createServer((sock) => this.onConn(sock)); + this.server.listen(0, '127.0.0.1', () => resolve(this.server.address().port)); + }); + } + onConn(sock) { + this.sock = sock; + let handshaken = false; + sock.on('data', (chunk) => { + this.buf = Buffer.concat([this.buf, chunk]); + if (!handshaken) { + const idx = this.buf.indexOf('\r\n\r\n'); + if (idx < 0) return; + const head = this.buf.subarray(0, idx).toString(); + this.buf = this.buf.subarray(idx + 4); + const key = /sec-websocket-key: *(.*)/i.exec(head)[1].trim(); + const proto = /sec-websocket-protocol: *(.*)/i.exec(head); + this.protocols = proto ? proto[1].trim() : ''; + const accept = crypto.createHash('sha1').update(key + '258EAFA5-E914-47DA-95CA-C5AB0DC85B11').digest('base64'); + sock.write('HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: ' + accept + '\r\n' + (proto ? 'Sec-WebSocket-Protocol: mqtt\r\n' : '') + '\r\n'); + handshaken = true; + } + this.parse(); + }); + sock.on('close', () => { this.closed = true; }); + sock.on('error', () => {}); + } + parse() { + while (this.buf.length >= 2) { + const b0 = this.buf[0], b1 = this.buf[1]; + let len = b1 & 0x7f, off = 2; + if (len === 126) { if (this.buf.length < 4) return; len = this.buf.readUInt16BE(2); off = 4; } + else if (len === 127) { if (this.buf.length < 10) return; len = Number(this.buf.readBigUInt64BE(2)); off = 10; } + const masked = (b1 & 0x80) !== 0; + const maskOff = off; + if (masked) off += 4; + if (this.buf.length < off + len) return; + const payload = Buffer.from(this.buf.subarray(off, off + len)); + if (masked) { const m = this.buf.subarray(maskOff, maskOff + 4); for (let i = 0; i < payload.length; i++) payload[i] ^= m[i % 4]; } + this.buf = this.buf.subarray(off + len); + const op = b0 & 0x0f; + this.opcodes.push(op); + if (op === 8) { this.closeReceived = true; try { this.sock.write(Buffer.from([0x88, 0])); this.sock.end(); } catch (e) {} } + else if (op === 0 || op === 1 || op === 2) { for (const b of payload) this.data.push(b); } + } + } + frame(op, payload) { + const p = Buffer.from(payload); + let hdr; + if (p.length < 126) hdr = Buffer.from([0x80 | op, p.length]); + else { hdr = Buffer.alloc(4); hdr[0] = 0x80 | op; hdr[1] = 126; hdr.writeUInt16BE(p.length, 2); } + this.sock.write(Buffer.concat([hdr, p])); + } + sendBinary(bytes) { this.frame(2, bytes); } + sendText(text) { this.frame(1, Buffer.from(text)); } + takeData() { const d = Uint8Array.from(this.data); this.data = []; return d; } + opcodeList() { return Uint8Array.from(this.opcodes); } + protocolHeader() { return this.protocols; } + isClosed() { return this.closed || !!this.closeReceived; } + stop() { try { if (this.sock) this.sock.destroy(); } catch (e) {} this.server.close(); } +} +"#)] +extern "C" { + type WsFake; + #[wasm_bindgen(constructor)] + fn new() -> WsFake; + #[wasm_bindgen(method)] + fn start(this: &WsFake) -> js_sys::Promise; + #[wasm_bindgen(method, js_name = sendBinary)] + fn send_binary(this: &WsFake, bytes: &[u8]); + #[wasm_bindgen(method, js_name = sendText)] + fn send_text(this: &WsFake, text: &str); + #[wasm_bindgen(method, js_name = takeData)] + fn take_data(this: &WsFake) -> Vec; + #[wasm_bindgen(method, js_name = opcodeList)] + fn opcode_list(this: &WsFake) -> Vec; + #[wasm_bindgen(method, js_name = protocolHeader)] + fn protocol_header(this: &WsFake) -> String; + #[wasm_bindgen(method, js_name = isClosed)] + fn is_closed(this: &WsFake) -> bool; + #[wasm_bindgen(method)] + fn stop(this: &WsFake); +} + +struct WsBroker { + server: WsFake, + inbox: RefCell>, +} + +impl WsBroker { + async fn next_packet(&self) -> Packet { + for _ in 0..400 { + self.inbox.borrow_mut().extend(self.server.take_data()); + if let Some(frame) = take_frame(&mut self.inbox.borrow_mut()) { + return frame.packet; + } + sleep(5).await; + } + panic!("client sent no packet over WebSocket"); + } + + async fn wait_closed(&self) -> bool { + for _ in 0..400 { + if self.server.is_closed() { + return true; + } + sleep(5).await; + } + false + } +} + +async fn connect_ws() -> (Rc, WsBroker) { + let server = WsFake::new(); + let port = JsFuture::from(server.start()) + .await + .unwrap() + .as_f64() + .unwrap(); + let broker = WsBroker { + server, + inbox: RefCell::new(Vec::new()), + }; + let client = Rc::new(WasmMqttClient::new("ws-client".to_string())); + let url = format!("ws://127.0.0.1:{port}/mqtt"); + let connecting = { + let client = Rc::clone(&client); + spawn_outcome(async move { client.connect(&url).await }) + }; + assert!(matches!(broker.next_packet().await, Packet::Connect(_))); + let connack = encode(&success()); + broker.server.send_binary(&connack[..2]); + sleep(10).await; + broker.server.send_binary(&connack[2..]); + settle(&connecting).await.expect("WebSocket connect failed"); + (client, broker) +} + +#[wasm_bindgen_test] +async fn mqtt_6_0_0_3_websocket_offers_mqtt_subprotocol() { + let (client, broker) = connect_ws().await; + let offered = broker.server.protocol_header(); + assert!( + offered.split(',').any(|p| p.trim() == "mqtt"), + "offered subprotocols: {offered:?}" + ); + client.disconnect().await.unwrap(); + broker.server.stop(); +} + +#[wasm_bindgen_test] +async fn mqtt_6_0_0_1_websocket_sends_binary_frames_only() { + let (client, broker) = connect_ws().await; + client.publish("a", b"payload").await.unwrap(); + assert!(matches!(broker.next_packet().await, Packet::Publish(_))); + let opcodes = broker.server.opcode_list(); + assert!(opcodes.iter().all(|op| *op == 2), "opcodes: {opcodes:?}"); + client.disconnect().await.unwrap(); + broker.server.stop(); +} + +#[wasm_bindgen_test] +async fn mqtt_6_0_0_1_websocket_text_frame_closes_connection() { + let (client, broker) = connect_ws().await; + broker.server.send_text("not mqtt"); + assert!( + broker.wait_closed().await, + "client kept the connection after a text frame" + ); + sleep(20).await; + assert!(!client.is_connected()); + broker.server.stop(); +} + +#[wasm_bindgen_test] +async fn mqtt_6_0_0_2_websocket_coalesced_and_split_packets() { + let (client, broker) = connect_ws().await; + let (callback, topics) = recorder(); + client.subscribe_with_callback("#", callback).await.unwrap(); + let id = match broker.next_packet().await { + Packet::Subscribe(subscribe) => subscribe.packet_id, + other => panic!("expected SUBSCRIBE, got {other:?}"), + }; + let mut coalesced = encode(&SubAckPacket::new(id).add_granted_qos(QoS::AtMostOnce)); + coalesced.extend([0xD0, 0x00]); + coalesced.extend(encode(&inbound_publish("one", QoS::AtMostOnce, None))); + let split = encode(&inbound_publish("two", QoS::AtMostOnce, None)); + coalesced.extend(&split[..1]); + broker.server.send_binary(&coalesced); + sleep(10).await; + broker.server.send_binary(&split[1..]); + wait_for_count(&topics, 2).await; + assert_eq!(topics.borrow().as_slice(), ["one", "two"]); + client.disconnect().await.unwrap(); + broker.server.stop(); +} diff --git a/crates/mqtt5/Cargo.toml b/crates/mqtt5/Cargo.toml index 0dfbc1bf..4c6d9ee5 100644 --- a/crates/mqtt5/Cargo.toml +++ b/crates/mqtt5/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "mqtt5" -version = "0.40.0" +version = "0.41.0" edition.workspace = true rust-version.workspace = true authors.workspace = true @@ -154,7 +154,7 @@ opentelemetry_sdk = { version = "0.32", features = ["trace", "rt-tokio", "metric opentelemetry-otlp = { version = "0.32", features = ["trace", "grpc-tonic", "metrics"], optional = true } tracing-opentelemetry = { version = "0.33", optional = true } opentelemetry = { version = "0.32", features = ["metrics", "trace"], optional = true } -mqtt5-protocol = "0.15.1" +mqtt5-protocol = "0.15.2" argon2 = { version = "0.5", optional = true } getrandom = "0.4.3" bebytes = { version = "3.0", features = ["bytes"] } diff --git a/crates/mqtt5/examples/deferred_ack.rs b/crates/mqtt5/examples/deferred_ack.rs index 4d0406df..d53ee13f 100644 --- a/crates/mqtt5/examples/deferred_ack.rs +++ b/crates/mqtt5/examples/deferred_ack.rs @@ -34,6 +34,7 @@ async fn main() -> Result<(), Box> { let options = ConnectOptions::new("deferred-ack-demo") .with_deferred_ack(true) .with_clean_start(false) + .with_resume_existing_session(true) .with_session_expiry_interval(3600) .with_receive_maximum(16); diff --git a/crates/mqtt5/src/broker/bridge/connection.rs b/crates/mqtt5/src/broker/bridge/connection.rs index e723dfe2..b417b07d 100644 --- a/crates/mqtt5/src/broker/bridge/connection.rs +++ b/crates/mqtt5/src/broker/bridge/connection.rs @@ -56,6 +56,7 @@ fn connect_options(config: &BridgeConfig) -> ConnectOptions { let mut options = ConnectOptions::new(&config.client_id) .with_deferred_ack(true) .with_clean_start(false) + .with_resume_existing_session(true) .with_session_expiry_interval(BRIDGE_SESSION_EXPIRY_SECS) .with_receive_maximum(receive_maximum); options.keep_alive = Duration::from_secs(u64::from(config.keepalive)); @@ -133,14 +134,12 @@ impl BridgeConnection { /// # Errors /// Returns an error if the configuration is invalid. pub fn new(config: BridgeConfig, router: Arc) -> Result { - // Validate configuration config .validate() .map_err(|e| BridgeError::ConfigurationError(e.to_string()))?; let client = Arc::new(MqttClient::with_options(connect_options(&config))); - // Create shutdown channel let (shutdown_tx, _) = broadcast::channel(1); Ok(Self { @@ -188,7 +187,9 @@ impl BridgeConnection { let _ = self.shutdown_tx.send(()); - let _ = self.client.disconnect().await; + if let Err(e) = self.client.disconnect().await { + debug!(bridge = %self.config.name, error = %e, "bridge client disconnect failed during stop"); + } let mut stats = self.stats.write().await; stats.connected = false; @@ -1009,10 +1010,8 @@ impl BridgeConnection { stats.current_broker = Some(address.to_string()); stats.on_primary = matches!(broker, ConnectedBroker::Primary); - // Store which broker we're connected to *self.current_broker.write().await = Some(broker); - // Flush any pending messages that were queued while disconnected self.flush_pending_messages().await; } @@ -1115,7 +1114,7 @@ impl BridgeConnection { async fn run_connection(&self) -> Result<()> { if !self.client.is_connected().await { self.register_ingress_callbacks().await?; - let _ = Box::pin(self.connect()).await?; + Box::pin(self.connect()).await?; self.setup_subscriptions().await?; } diff --git a/crates/mqtt5/src/client/direct/ack.rs b/crates/mqtt5/src/client/direct/ack.rs index 6d42e2fa..adce6b28 100644 --- a/crates/mqtt5/src/client/direct/ack.rs +++ b/crates/mqtt5/src/client/direct/ack.rs @@ -1,5 +1,5 @@ use crate::callback::panic_message; -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; use std::panic::{catch_unwind, AssertUnwindSafe}; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, OnceLock}; @@ -9,7 +9,9 @@ use tokio::sync::mpsc; use tracing::{debug, warn}; use crate::callback::CallbackId; +use crate::client::direct::handlers::ack_fits_server_maximum; use crate::client::direct::unified::UnifiedWriter; +use crate::packet::is_valid_publish_ack_reason_code; use crate::packet::puback::PubAckPacket; use crate::packet::publish::PublishPacket; use crate::packet::pubrec::PubRecPacket; @@ -33,14 +35,59 @@ pub(crate) enum AckKind { Ack, Reject(ReasonCode), DropAuto, + Automatic { + packet: Packet, + release_inbound: bool, + }, } pub(crate) struct AckRequest { + seq: Option, packet_id: u16, qos: QoS, kind: AckKind, } +#[derive(Default)] +struct AckOrder { + head: u64, + slots: VecDeque>, +} + +impl AckOrder { + fn reserve(&mut self) -> u64 { + self.slots.push_back(None); + self.head + self.slots.len() as u64 - 1 + } + + fn discard(&mut self) { + self.head += self.slots.len() as u64; + self.slots.clear(); + } + + fn release(&mut self, request: AckRequest) -> Vec { + if request.seq.is_some_and(|seq| seq < self.head) { + return Vec::new(); + } + let index = request + .seq + .and_then(|seq| seq.checked_sub(self.head)) + .and_then(|offset| usize::try_from(offset).ok()); + let Some(slot) = index.and_then(|i| self.slots.get_mut(i)) else { + return vec![request]; + }; + *slot = Some(request); + let mut ready = Vec::new(); + while self.slots.front().is_some_and(Option::is_some) { + if let Some(Some(next)) = self.slots.pop_front() { + ready.push(next); + } + self.head += 1; + } + ready + } +} + /// A capability to acknowledge exactly one inbound `QoS` > 0 message after the /// application has durably processed it. /// @@ -50,6 +97,7 @@ pub(crate) struct AckRequest { /// Dropping it without resolving emits a non-success acknowledgement and warns, so a /// forgotten token never wedges the window (`DeferredAckToken.tla`, obligation 7). pub struct AckToken { + seq: Option, packet_id: u16, qos: QoS, armed: bool, @@ -83,10 +131,12 @@ impl AckToken { /// broker replays the message on reconnect, your callback will see it again. Reject /// must therefore be idempotent, like the rest of deferred delivery. /// - /// A non-error `reason` (Reason Code below `0x80`, e.g. `Success`) is normalized to - /// [`ReasonCode::UnspecifiedError`], so `reject` can never behave as an [`ack`](Self::ack). + /// A `reason` that is not an error PUBACK/PUBREC Reason Code (e.g. `Success`, or a + /// code such as `ServerBusy` that those packets do not allow) is normalized to + /// [`ReasonCode::UnspecifiedError`], so `reject` can never behave as an [`ack`](Self::ack) + /// or put a code on the wire that `[MQTT-3.4.2-1]`/`[MQTT-3.5.2-1]` forbid. pub fn reject(mut self, reason: ReasonCode) { - let reason = if reason.is_error() { + let reason = if reason.is_error() && is_valid_publish_ack_reason_code(reason) { reason } else { ReasonCode::UnspecifiedError @@ -100,6 +150,7 @@ impl AckToken { } self.armed = false; let _ = self.sender.send(AckRequest { + seq: self.seq, packet_id: self.packet_id, qos: self.qos, kind, @@ -130,6 +181,7 @@ impl Drop for AckToken { /// tokens only ever enqueue an `AckRequest` here. pub(crate) struct AckDispatcher { tx: mpsc::UnboundedSender, + order: Arc>, writer_slot: WriterSlot, session: Arc>, pending_rx: tokio::sync::Mutex>>, @@ -140,6 +192,7 @@ impl AckDispatcher { let (tx, rx) = mpsc::unbounded_channel::(); Self { tx, + order: Arc::new(Mutex::new(AckOrder::default())), writer_slot: Arc::new(tokio::sync::Mutex::new(None)), session, pending_rx: tokio::sync::Mutex::new(Some(rx)), @@ -154,16 +207,25 @@ impl AckDispatcher { }; let slot = Arc::clone(&self.writer_slot); let session = Arc::clone(&self.session); + let order = Arc::clone(&self.order); tokio::spawn(async move { while let Some(request) = rx.recv().await { - Self::handle(&request, &slot, &session).await; + let ready = order.lock().release(request); + for next in ready { + Self::handle(next, &slot, &session).await; + } } }); } + fn reserve(&self, qos: QoS) -> Option { + (qos != QoS::AtMostOnce).then(|| self.order.lock().reserve()) + } + /// Mints a token for a delivered inbound message. pub(crate) fn token(&self, packet_id: u16, qos: QoS) -> AckToken { AckToken { + seq: self.reserve(qos), packet_id, qos, armed: true, @@ -188,10 +250,15 @@ impl AckDispatcher { *self.writer_slot.lock().await = None; } + pub(crate) fn discard_pending(&self) { + self.order.lock().discard(); + } + /// Re-sends an acknowledgement for a duplicate that was already resolved, /// without a token (used on a post-reconnect replay). pub(crate) fn enqueue(&self, packet_id: u16, qos: QoS, kind: AckKind) { let _ = self.tx.send(AckRequest { + seq: self.reserve(qos), packet_id, qos, kind, @@ -206,7 +273,7 @@ impl AckDispatcher { /// the correct replay on reconnect. For a `QoS` 2 error acknowledgement the per-id deferred /// state is cleared, per `[MQTT-4.3.3-9]` (a later same-id PUBLISH is a new message). async fn handle( - request: &AckRequest, + request: AckRequest, slot: &WriterSlot, session: &Arc>, ) { @@ -214,6 +281,20 @@ impl AckDispatcher { AckKind::Ack => ReasonCode::Success, AckKind::Reject(r) => r, AckKind::DropAuto => DROP_REASON, + AckKind::Automatic { + packet, + release_inbound, + } => { + Self::write(request.packet_id, packet, slot, session).await; + if release_inbound { + session + .read() + .await + .acknowledge_inbound(request.packet_id) + .await; + } + return; + } }; let packet = match request.qos { QoS::AtMostOnce => return, @@ -247,6 +328,18 @@ impl AckDispatcher { } } + Self::write(request.packet_id, packet, slot, session).await; + } + + async fn write( + packet_id: u16, + packet: Packet, + slot: &WriterSlot, + session: &Arc>, + ) { + if !ack_fits_server_maximum(session, &packet).await { + return; + } let writer = slot.lock().await.clone(); let written = match &writer { Some(handle) => handle.lock().await.write_packet(packet).await.is_ok(), @@ -254,8 +347,8 @@ impl AckDispatcher { }; if !written { debug!( - packet_id = request.packet_id, - "Deferred ack not written (disconnected); resolution recorded for replay" + packet_id, + "Ack not written (disconnected); resolution recorded for replay" ); } } @@ -390,13 +483,63 @@ impl AckCallbackManager { #[cfg(test)] mod tests { - use super::{AckCallbackManager, AckDispatcher, AckPublishCallback}; + use super::{ + AckCallbackManager, AckDispatcher, AckKind, AckOrder, AckPublishCallback, AckRequest, + }; use crate::packet::publish::PublishPacket; use crate::session::state::{SessionConfig, SessionState}; use crate::QoS; use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::Arc; + fn request(seq: Option, packet_id: u16) -> AckRequest { + AckRequest { + seq, + packet_id, + qos: QoS::AtLeastOnce, + kind: AckKind::Ack, + } + } + + fn released_ids(order: &mut AckOrder, request: AckRequest) -> Vec { + order + .release(request) + .iter() + .map(|released| released.packet_id) + .collect() + } + + #[test] + fn ack_order_holds_later_acks_until_earlier_ones_resolve() { + let mut order = AckOrder::default(); + let first = order.reserve(); + let second = order.reserve(); + let third = order.reserve(); + + assert!(released_ids(&mut order, request(Some(third), 3)).is_empty()); + assert!(released_ids(&mut order, request(Some(second), 2)).is_empty()); + assert_eq!( + released_ids(&mut order, request(Some(first), 1)), + vec![1, 2, 3] + ); + assert_eq!(released_ids(&mut order, request(None, 9)), vec![9]); + + let fourth = order.reserve(); + assert_eq!(released_ids(&mut order, request(Some(fourth), 4)), vec![4]); + } + + #[test] + fn discarded_acks_and_stale_tokens_are_never_released() { + let mut order = AckOrder::default(); + let held = order.reserve(); + let queued = order.reserve(); + assert!(released_ids(&mut order, request(Some(queued), 2)).is_empty()); + order.discard(); + assert!(released_ids(&mut order, request(Some(held), 1)).is_empty()); + let fresh = order.reserve(); + assert_eq!(released_ids(&mut order, request(Some(fresh), 3)), vec![3]); + } + #[tokio::test] async fn panicking_ack_callback_does_not_stop_later_delivery() { let mgr = AckCallbackManager::new(); diff --git a/crates/mqtt5/src/client/direct/handlers.rs b/crates/mqtt5/src/client/direct/handlers.rs index 529822e0..c32dcba1 100644 --- a/crates/mqtt5/src/client/direct/handlers.rs +++ b/crates/mqtt5/src/client/direct/handlers.rs @@ -7,7 +7,8 @@ use crate::packet::publish::PublishPacket; use crate::packet::Packet; use crate::protocol::v5::properties::Properties; use crate::session::state::AckResolution; -use crate::session::SessionState; +use crate::session::{SessionState, TopicAliasManager}; +use crate::transport::packet_io::encode_packet_to_buffer; use crate::transport::PacketWriter; use crate::QoS; use parking_lot::Mutex; @@ -30,6 +31,7 @@ pub(super) struct AckDelivery<'a> { /// reader context so the entry point stays a small number of arguments. pub(super) struct IncomingHandlers<'a> { pub(super) session: &'a Arc>, + pub(super) topic_aliases: &'a Mutex, pub(super) callback_manager: &'a Arc, pub(super) keepalive_state: &'a Arc>, pub(super) codec_registry: Option<&'a Arc>, @@ -44,7 +46,8 @@ pub(super) async fn handle_incoming_packet_with_writer( ) -> Result<()> { let session = handlers.session; match packet { - Packet::Publish(publish) => { + Packet::Publish(mut publish) => { + validate_inbound_publish(&mut publish, handlers.topic_aliases)?; handle_publish_with_ack( publish, writer, @@ -72,6 +75,74 @@ pub(super) async fn handle_incoming_packet_with_writer( } } +enum AckRoute<'a> { + Direct(&'a Arc>), + Ordered(&'a Arc), +} + +fn validate_inbound_publish( + publish: &mut PublishPacket, + topic_aliases: &Mutex, +) -> Result<()> { + if publish.properties.subscription_identifiers().contains(&0) { + return Err(MqttError::ProtocolError( + "inbound PUBLISH carries Subscription Identifier 0".to_string(), + )); + } + let mut aliases = topic_aliases.lock(); + match publish.properties.get_topic_alias() { + Some(alias) if alias == 0 || alias > aliases.topic_alias_maximum() => { + Err(MqttError::TopicAliasInvalid(alias)) + } + Some(alias) if publish.topic_name.is_empty() => { + let topic = aliases.get_topic(alias).ok_or_else(|| { + MqttError::ProtocolError(format!("Topic Alias {alias} has no mapping")) + })?; + publish.topic_name = topic.to_string(); + Ok(()) + } + Some(alias) => aliases.register_alias(alias, &publish.topic_name), + None if publish.topic_name.is_empty() => Err(MqttError::ProtocolError( + "zero-length Topic Name without a Topic Alias".to_string(), + )), + None => Ok(()), + } +} + +pub(super) async fn ack_fits_server_maximum( + session: &Arc>, + packet: &Packet, +) -> bool { + let mut buf = bytes::BytesMut::new(); + let fits = encode_packet_to_buffer(packet, &mut buf).is_ok() + && session + .read() + .await + .check_packet_size(buf.len()) + .await + .is_ok(); + if !fits { + tracing::warn!( + packet = packet.packet_type_name(), + size = buf.len(), + "Acknowledgement exceeds the server Maximum Packet Size; discarded [MQTT-3.2.2-15]" + ); + } + fits +} + +async fn write_ack( + writer: &Arc>, + session: &Arc>, + packet: Packet, +) -> Result<()> { + if ack_fits_server_maximum(session, &packet).await { + writer.lock().await.write_packet(packet).await + } else { + Ok(()) + } +} + pub(super) async fn handle_publish_with_ack( mut publish: crate::packet::publish::PublishPacket, writer: &Arc>, @@ -93,17 +164,22 @@ pub(super) async fn handle_publish_with_ack( } } + let route = match ack_delivery { + Some(ack) if flow_id.is_none() => AckRoute::Ordered(ack.dispatcher), + _ => AckRoute::Direct(writer), + }; + let already_delivered = match publish.qos { crate::QoS::AtMostOnce => false, crate::QoS::AtLeastOnce => { if let Some(packet_id) = publish.packet_id { - ack_qos1_inbound(packet_id, writer, session, flow_id).await?; + ack_qos1_inbound(packet_id, &route, session, flow_id).await?; } false } crate::QoS::ExactlyOnce => { if let Some(packet_id) = publish.packet_id { - let receipt = ack_qos2_inbound(packet_id, writer, session, flow_id).await?; + let receipt = ack_qos2_inbound(packet_id, &route, session, flow_id).await?; receipt == Qos2Receipt::Duplicate } else { false @@ -204,7 +280,7 @@ async fn resend_matching_ack( async fn ack_qos1_inbound( packet_id: u16, - writer: &Arc>, + route: &AckRoute<'_>, session: &Arc>, flow_id: Option, ) -> Result<()> { @@ -225,16 +301,26 @@ async fn ack_qos1_inbound( .await; } - let puback = crate::packet::puback::PubAckPacket { + let puback = Packet::PubAck(crate::packet::puback::PubAckPacket { packet_id, reason_code: crate::protocol::v5::reason_codes::ReasonCode::Success, properties: Properties::default(), + }); + let writer = match route { + AckRoute::Direct(writer) => writer, + AckRoute::Ordered(dispatcher) => { + dispatcher.enqueue( + packet_id, + QoS::AtLeastOnce, + AckKind::Automatic { + packet: puback, + release_inbound: true, + }, + ); + return Ok(()); + } }; - writer - .lock() - .await - .write_packet(Packet::PubAck(puback)) - .await?; + write_ack(writer, session, puback).await?; session .read() @@ -256,7 +342,7 @@ enum Qos2Receipt { async fn ack_qos2_inbound( packet_id: u16, - writer: &Arc>, + route: &AckRoute<'_>, session: &Arc>, flow_id: Option, ) -> Result { @@ -279,21 +365,28 @@ async fn ack_qos2_inbound( let first_receipt = session.write().await.mark_pubrec_pending(packet_id).await; - let pubrec = crate::packet::pubrec::PubRecPacket { + let pubrec = Packet::PubRec(crate::packet::pubrec::PubRecPacket { packet_id, reason_code: crate::protocol::v5::reason_codes::ReasonCode::Success, properties: Properties::default(), - }; - if let Err(e) = writer - .lock() - .await - .write_packet(Packet::PubRec(pubrec)) - .await - { - if first_receipt { - session.write().await.remove_pubrec(packet_id).await; + }); + match route { + AckRoute::Direct(writer) => { + if let Err(e) = write_ack(writer, session, pubrec).await { + if first_receipt { + session.write().await.remove_pubrec(packet_id).await; + } + return Err(e); + } } - return Err(e); + AckRoute::Ordered(dispatcher) => dispatcher.enqueue( + packet_id, + QoS::ExactlyOnce, + AckKind::Automatic { + packet: pubrec, + release_inbound: false, + }, + ), } if first_receipt { @@ -401,11 +494,7 @@ pub(super) async fn handle_pubrel( properties: Properties::default(), }; - writer - .lock() - .await - .write_packet(Packet::PubComp(pubcomp)) - .await?; + write_ack(writer, session, Packet::PubComp(pubcomp)).await?; session .read() diff --git a/crates/mqtt5/src/client/direct/keepalive.rs b/crates/mqtt5/src/client/direct/keepalive.rs index c5fe9aed..e239dd82 100644 --- a/crates/mqtt5/src/client/direct/keepalive.rs +++ b/crates/mqtt5/src/client/direct/keepalive.rs @@ -149,6 +149,8 @@ pub(super) async fn keepalive_task_with_writer( let timed_out = keepalive_state.lock().is_timeout(timeout_duration); if timed_out { tracing::error!("Keepalive timeout - no PINGRESP received"); + super::reader::close_connection(&writer, &crate::error::MqttError::KeepAliveTimeout) + .await; lifecycle.end(DisconnectReason::KeepAliveTimeout).await; break; } diff --git a/crates/mqtt5/src/client/direct/mod.rs b/crates/mqtt5/src/client/direct/mod.rs index 5ad8e952..7687896e 100644 --- a/crates/mqtt5/src/client/direct/mod.rs +++ b/crates/mqtt5/src/client/direct/mod.rs @@ -5,14 +5,16 @@ pub(crate) mod ack; mod handlers; mod keepalive; +mod outbound; mod reader; +mod replay; mod unified; pub use ack::AckToken; pub(crate) use ack::{AckCallbackManager, AckDispatcher}; use parking_lot::Mutex; -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::Arc; use tokio::sync::oneshot; @@ -31,8 +33,10 @@ use crate::packet::unsuback::UnsubAckPacket; use crate::packet::unsubscribe::UnsubscribePacket; use crate::packet::{MqttPacket, Packet}; use crate::packet_id::PacketIdGenerator; -use crate::protocol::v5::properties::Properties; +use crate::protocol::v5::properties::{Properties, PropertyId, PropertyValue}; use crate::protocol::v5::reason_codes::ReasonCode; +use crate::session::flow_control::FlowControlManager; +use crate::session::state::OutboundReplay; use crate::session::subscription::Subscription; use crate::session::SessionState; use crate::transport::{PacketIo, PacketWriter, TransportType}; @@ -60,6 +64,7 @@ use keepalive::{keepalive_task_with_writer, KeepaliveState}; #[cfg(feature = "transport-quic")] use reader::quic_stream_acceptor_task; use reader::{packet_reader_task_with_responses, PacketReaderContext}; +use replay::{PublishPolicy, SessionReplay}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum AutomaticReconnectLifecycle { @@ -76,6 +81,37 @@ pub(crate) enum SubscriptionPersistence { pub(crate) type StoredSubscription = (String, SubscriptionOptions, Option, CallbackId); pub(crate) type StoredSubscriptions = Arc>>; pub(crate) type ConnectionEpoch = Arc; +pub(crate) type PendingAcks = Arc>>>; + +#[derive(Debug)] +pub(crate) enum StagedPublish { + Queued(PublishResult), + Ready(PublishPacket), +} + +pub(crate) struct PublishAck { + rx: oneshot::Receiver, + packet_id: u16, + pending: PendingAcks, +} + +impl PublishAck { + pub(crate) async fn wait(self) -> Result<()> { + match tokio::time::timeout(Duration::from_secs(10), self.rx).await { + Ok(Ok(reason_code)) if reason_code.is_error() => { + Err(MqttError::PublishFailed(reason_code)) + } + Ok(Ok(_)) => Ok(()), + Ok(Err(_)) => Err(MqttError::ProtocolError( + "Acknowledgment channel closed".to_string(), + )), + Err(_) => { + self.pending.lock().remove(&self.packet_id); + Err(MqttError::Timeout) + } + } + } +} pub struct DirectClientInner { pub writer: Option>>, @@ -109,21 +145,23 @@ pub struct DirectClientInner { pub packet_id_generator: PacketIdGenerator, pub pending_subacks: Arc>>>, pub pending_unsubacks: Arc>>>, - pub pending_pubacks: Arc>>>, - pub pending_pubcomps: Arc>>>, + pub pending_pubacks: PendingAcks, + pub pending_pubcomps: PendingAcks, pub reconnect_attempt: u32, pub last_address: Option, pub automatic_reconnect_lifecycle: AutomaticReconnectLifecycle, pub server_redirect: Option, - pub queued_messages: Arc>>, + pub queued_messages: Arc>>, pub stored_subscriptions: StoredSubscriptions, pub stored_ack_subscriptions: StoredSubscriptions, pub queue_on_disconnect: bool, pub server_max_qos: Arc>>, + pub server_retain_available: Arc, pub auth_handler: Option>, pub auth_method: Option, pub keepalive_state: Arc>, pub negotiated_keep_alive_secs: AtomicU64, + server_capabilities: outbound::ServerCapabilities, #[cfg(feature = "transport-quic")] pub cached_quic_client_config: Option, #[cfg(feature = "transport-quic")] @@ -179,15 +217,17 @@ impl DirectClientInner { last_address: None, automatic_reconnect_lifecycle: AutomaticReconnectLifecycle::Armed, server_redirect: None, - queued_messages: Arc::new(Mutex::new(Vec::new())), + queued_messages: Arc::new(Mutex::new(VecDeque::new())), stored_subscriptions: Arc::new(Mutex::new(Vec::new())), stored_ack_subscriptions: Arc::new(Mutex::new(Vec::new())), queue_on_disconnect, server_max_qos: Arc::new(Mutex::new(None)), + server_retain_available: Arc::new(AtomicBool::new(true)), auth_handler: None, auth_method, keepalive_state: Arc::new(Mutex::new(KeepaliveState::default())), negotiated_keep_alive_secs: AtomicU64::new(initial_keep_alive_secs), + server_capabilities: outbound::ServerCapabilities::default(), #[cfg(feature = "transport-quic")] cached_quic_client_config: None, #[cfg(feature = "transport-quic")] @@ -289,19 +329,6 @@ impl DirectClientInner { pub fn set_queue_on_disconnect(&mut self, enabled: bool) { self.queue_on_disconnect = enabled; } - - /// # Errors - /// - /// Returns an error if the operation fails - pub async fn send_packet(&mut self, packet: Packet) -> Result<()> { - if let Some(writer) = &self.writer { - let mut writer_guard = writer.lock().await; - writer_guard.write_packet(packet).await?; - Ok(()) - } else { - Err(MqttError::NotConnected) - } - } } impl DirectClientInner { @@ -314,6 +341,11 @@ impl DirectClientInner { match auth.reason_code { ReasonCode::ContinueAuthentication => { + let method = self.auth_method.clone().ok_or_else(|| { + MqttError::ProtocolError( + "AUTH received but CONNECT carried no Authentication Method".to_string(), + ) + })?; let handler = self .auth_handler .as_ref() @@ -326,7 +358,6 @@ impl DirectClientInner { match response { AuthResponse::Continue(data) => { - let method = self.auth_method.clone().unwrap_or_default(); let auth_packet = AuthPacket::continue_authentication(method, Some(data))?; transport.write_packet(Packet::Auth(auth_packet)).await?; } @@ -411,16 +442,14 @@ impl DirectClientInner { return Err(MqttError::ConnectionRefused(connack.reason_code)); } - if let Some(max_qos) = connack.properties.get_maximum_qos() { - *self.server_max_qos.lock() = Some(max_qos); - tracing::debug!("Server maximum QoS: {}", max_qos); - } else { - *self.server_max_qos.lock() = None; + if connack.session_present && !self.holds_session_state() { + return Err(Self::reject_unexpected_session_present(&mut transport).await); } - self.apply_negotiated_keep_alive(connack.properties.get_server_keep_alive()); - - self.apply_negotiated_packet_sizes(&connack).await?; + let receive_maximum = self.apply_server_capabilities(&connack).await?; + self.apply_negotiated_capabilities(&connack).await; + let replay_items = self.session.read().await.outbound_replay().await; + let replay_slots = self.reset_send_quota(receive_maximum, &replay_items).await; let protocol_version = self.options.protocol_version.as_u8(); let (reader, writer) = match transport { @@ -471,11 +500,13 @@ impl DirectClientInner { } }; + let reader = reader.with_maximum_packet_size(self.options.properties.maximum_packet_size); let connection_epoch = self.advance_connection_epoch(); let writer_arc = Arc::new(tokio::sync::Mutex::new(writer)); self.ack_dispatcher .set_writer(Arc::clone(&writer_arc)) .await; + let replay_writer = Arc::downgrade(&writer_arc); self.writer = Some(writer_arc); self.set_connected(true); @@ -483,34 +514,153 @@ impl DirectClientInner { self.start_background_tasks(reader, connection_epoch)?; tracing::debug!("Background tasks started successfully"); + if let Some(slots) = replay_slots { + let replay = SessionReplay { + items: replay_items, + slots, + session: Arc::clone(&self.session), + writer: replay_writer, + queued: Arc::clone(&self.queued_messages), + policy: self.publish_policy(), + }; + tokio::spawn(replay.run()); + } + Ok(ConnectResult { session_present: connack.session_present, }) } + fn holds_session_state(&self) -> bool { + !self.options.clean_start + && (self.connection_epoch.load(Ordering::SeqCst) > 0 + || self.options.resume_existing_session) + } + + async fn apply_server_capabilities( + &mut self, + connack: &crate::packet::connack::ConnAckPacket, + ) -> Result { + if !connack.session_present { + self.discard_session_state().await; + } + self.adopt_assigned_client_identifier(connack).await; + + if let Some(max_qos) = connack.properties.get_maximum_qos() { + *self.server_max_qos.lock() = Some(max_qos); + tracing::debug!("Server maximum QoS: {}", max_qos); + } else { + *self.server_max_qos.lock() = None; + } + self.server_retain_available.store( + !matches!( + connack.properties.get(PropertyId::RetainAvailable), + Some(PropertyValue::Byte(0)) + ), + Ordering::SeqCst, + ); + + self.apply_negotiated_keep_alive(connack.properties.get_server_keep_alive()); + + self.apply_negotiated_packet_sizes(connack).await + } + + async fn reject_unexpected_session_present(transport: &mut TransportType) -> MqttError { + tracing::warn!( + "CONNACK reported Session Present=1 but the client holds no session state; closing the connection" + ); + let disconnect = crate::packet::disconnect::DisconnectPacket { + reason_code: ReasonCode::ProtocolError, + properties: Properties::default(), + }; + if let Err(e) = transport.write_packet(Packet::Disconnect(disconnect)).await { + tracing::debug!("Failed to send DISCONNECT for unexpected Session Present: {e}"); + } + MqttError::ProtocolError( + "CONNACK Session Present=1 but the client holds no session state".to_string(), + ) + } + + async fn discard_session_state(&self) { + self.ack_dispatcher.discard_pending(); + let session = self.session.read().await; + session.discard_outbound_state().await; + session.flow_control().read().await.clear_inbound().await; + if session.clear_all_inbound_state().await { + if self.options.deferred_ack { + tracing::warn!( + "Reconnected with session_present=0; cleared stale inbound QoS 2 \ + de-duplication state. Any outstanding AckTokens are now stale because the \ + broker no longer holds the session that delivered their messages." + ); + } else { + tracing::debug!( + "Reconnected with session_present=0; cleared stale inbound QoS 2 de-duplication state" + ); + } + } + } + + async fn adopt_assigned_client_identifier( + &mut self, + connack: &crate::packet::connack::ConnAckPacket, + ) { + if let Some(PropertyValue::Utf8String(assigned)) = + connack.properties.get(PropertyId::AssignedClientIdentifier) + { + tracing::debug!(client_id = %assigned, "Adopting server assigned client identifier"); + self.session.write().await.set_client_id(assigned.clone()); + self.options.client_id.clone_from(assigned); + } + } + + async fn reset_send_quota( + &self, + receive_maximum: u16, + replay_items: &[OutboundReplay], + ) -> Option> { + let retained_in_flight: Vec = replay_items + .iter() + .filter_map(|item| match item { + OutboundReplay::PubRel(packet_id) => Some(*packet_id), + OutboundReplay::Publish(_) => None, + }) + .collect(); + let replay = !replay_items.is_empty() || !self.queued_messages.lock().is_empty(); + let flow = Arc::clone(self.session.read().await.flow_control()); + let mut flow = flow.write().await; + flow.reset_for_connection(receive_maximum, &retained_in_flight, replay) + .await + } + + fn publish_policy(&self) -> PublishPolicy { + PublishPolicy { + maximum_qos: *self.server_max_qos.lock(), + retain_available: self.server_retain_available.load(Ordering::SeqCst), + } + } + async fn apply_negotiated_packet_sizes( &self, connack: &crate::packet::connack::ConnAckPacket, - ) -> Result<()> { + ) -> Result { let session = self.session.write().await; - match connack.properties.get_receive_maximum() { + let receive_maximum = match connack.properties.get_receive_maximum() { Some(0) => { return Err(MqttError::ProtocolError( "server advertised a Receive Maximum of 0".to_string(), )); } Some(server_receive_maximum) => { - session.set_receive_maximum(server_receive_maximum).await; tracing::debug!("Server Receive Maximum: {}", server_receive_maximum); + server_receive_maximum } - None => session.set_receive_maximum(65535).await, - } + None => 65535, + }; - if self.options.deferred_ack { - if let Some(receive_maximum) = self.options.properties.receive_maximum { - session.set_inbound_receive_maximum(receive_maximum).await; - } + if let Some(receive_maximum) = self.options.properties.receive_maximum { + session.set_inbound_receive_maximum(receive_maximum).await; } if let Some(max_packet_size) = self.options.properties.maximum_packet_size { @@ -529,7 +679,19 @@ impl DirectClientInner { None => session.reset_server_maximum_packet_size().await, } - Ok(()) + Ok(receive_maximum) + } + + async fn apply_negotiated_capabilities( + &mut self, + connack: &crate::packet::connack::ConnAckPacket, + ) { + self.server_capabilities = outbound::ServerCapabilities::from_connack(connack); + self.session + .read() + .await + .set_topic_alias_maximum_out(connack.topic_alias_maximum().unwrap_or(0)) + .await; } /// # Errors @@ -582,21 +744,26 @@ impl DirectClientInner { return Err(MqttError::NotConnected); } - if send_disconnect { - if let Some(ref writer) = self.writer { - let disconnect = crate::packet::disconnect::DisconnectPacket { - reason_code: crate::protocol::v5::reason_codes::ReasonCode::Success, - properties: crate::protocol::v5::properties::Properties::default(), - }; - let _ = writer - .lock() - .await - .write_packet(Packet::Disconnect(disconnect)) - .await; + self.set_connected(false); + if let Some(ref writer) = self.writer { + let disconnect = send_disconnect.then(|| { + Packet::Disconnect(crate::packet::disconnect::DisconnectPacket::new( + ReasonCode::Success, + )) + }); + if let Err(e) = writer.lock().await.close(disconnect).await { + tracing::debug!("Closing network connection on disconnect: {e}"); } } self.reset_connection_runtime(b"disconnect").await; + self.session + .read() + .await + .flow_control() + .read() + .await + .close_send_quota(); Ok(()) } @@ -614,26 +781,60 @@ impl DirectClientInner { payload: Vec, options: &PublishOptions, ) -> Result { - let mut publish = PublishPacket { - topic_name: topic, - packet_id: Some(SIZE_PROBE_PACKET_ID), - payload: payload.into(), - qos: options.qos, - retain: options.retain, - dup: false, - properties: options.properties.clone().into(), - protocol_version: self.options.protocol_version.as_u8(), - stream_id: None, - }; + let mut publish = self + .with_aliased_topic(PublishPacket { + topic_name: topic, + packet_id: Some(SIZE_PROBE_PACKET_ID), + payload: payload.into(), + qos: options.qos, + retain: options.retain, + dup: false, + properties: options.properties.clone().into(), + protocol_version: self.options.protocol_version.as_u8(), + stream_id: None, + }) + .await?; self.check_publish_size(&publish).await?; - let packet_id = self.packet_id_generator.next(); + let packet_id = self.allocate_packet_id().await?; publish.packet_id = Some(packet_id); - self.queued_messages.lock().push(publish); + self.queued_messages.lock().push_back(publish); Ok(PublishResult::QoS1Or2 { packet_id }) } + async fn with_aliased_topic(&self, mut publish: PublishPacket) -> Result { + if let Some(alias) = publish + .topic_alias() + .filter(|_| publish.topic_name.is_empty()) + { + let session = self.session.read().await; + let aliases = session.topic_alias_out().read().await; + publish.topic_name = aliases + .get_topic(alias) + .map(str::to_string) + .ok_or(MqttError::TopicAliasInvalid(alias))?; + } + Ok(publish) + } + + async fn allocate_packet_id(&self) -> Result { + self.session + .read() + .await + .allocate_packet_id(&self.packet_id_generator, |packet_id| { + self.pending_subacks.lock().contains_key(&packet_id) + || self.pending_unsubacks.lock().contains_key(&packet_id) + || self + .queued_messages + .lock() + .iter() + .any(|queued| queued.packet_id == Some(packet_id)) + }) + .await + .ok_or(MqttError::PacketIdExhausted) + } + /// # Errors /// /// Returns `PacketTooLarge` if the encoded packet exceeds the negotiated @@ -644,62 +845,31 @@ impl DirectClientInner { self.session.read().await.check_packet_size(buf.len()).await } - fn setup_publish_acknowledgment( - &self, - qos: QoS, - packet_id: Option, - ) -> Option> { - match qos { - QoS::AtMostOnce => None, - QoS::AtLeastOnce => { - let (tx, rx) = oneshot::channel(); - if let Some(pid) = packet_id { - self.pending_pubacks.lock().insert(pid, tx); - } - Some(rx) - } - QoS::ExactlyOnce => { - let (tx, rx) = oneshot::channel(); - if let Some(pid) = packet_id { - self.pending_pubcomps.lock().insert(pid, tx); - } - Some(rx) - } - } + async fn check_packet_fits(&self, packet: &impl MqttPacket) -> Result<()> { + let mut buf = bytes::BytesMut::new(); + packet.encode(&mut buf)?; + self.session.read().await.check_packet_size(buf.len()).await } - async fn wait_for_acknowledgment( - &self, - rx: oneshot::Receiver, - qos: QoS, - packet_id: Option, - ) -> Result<()> { - let timeout = Duration::from_secs(10); - match tokio::time::timeout(timeout, rx).await { - Ok(Ok(reason_code)) => { - if reason_code.is_error() { - return Err(MqttError::PublishFailed(reason_code)); - } - Ok(()) - } - Ok(Err(_)) => Err(MqttError::ProtocolError( - "Acknowledgment channel closed".to_string(), - )), - Err(_) => { - if let Some(pid) = packet_id { - match qos { - QoS::AtLeastOnce => { - self.pending_pubacks.lock().remove(&pid); - } - QoS::ExactlyOnce => { - self.pending_pubcomps.lock().remove(&pid); - } - QoS::AtMostOnce => {} - } - } - Err(MqttError::Timeout) - } - } + pub(crate) async fn check_unsubscribe(&self, packet: &UnsubscribePacket) -> Result<()> { + outbound::check_unsubscribe(packet)?; + self.check_packet_fits(packet).await + } + + fn setup_publish_acknowledgment(&self, qos: QoS, packet_id: Option) -> Option { + let pending = match qos { + QoS::AtMostOnce => return None, + QoS::AtLeastOnce => &self.pending_pubacks, + QoS::ExactlyOnce => &self.pending_pubcomps, + }; + let packet_id = packet_id?; + let (tx, rx) = oneshot::channel(); + pending.lock().insert(packet_id, tx); + Some(PublishAck { + rx, + packet_id, + pending: Arc::clone(pending), + }) } pub(super) async fn release_outbound_quota( @@ -707,26 +877,38 @@ impl DirectClientInner { packet_id: Option, ) { if let Some(pid) = packet_id { - let flow = session.read().await.flow_control().clone(); - let _ = flow.read().await.acknowledge(pid).await; + let session = session.read().await; + session.complete_outbound(pid).await; + let flow = Arc::clone(session.flow_control()); + drop(session); + Self::release_send_quota(&flow, pid).await; } } - /// # Errors - /// - /// Returns an error if the operation fails - pub async fn publish( + async fn release_send_quota( + flow: &Arc>, + packet_id: u16, + ) { + if let Err(e) = flow.read().await.acknowledge(packet_id).await { + tracing::trace!(packet_id, "No send quota held: {e}"); + } + } + + pub(crate) async fn stage_publish( &self, topic: String, payload: Vec, options: PublishOptions, - ) -> Result { + ) -> Result { + outbound::check_publish(&topic, &options)?; + if !self.is_connected() && self.queue_on_disconnect && options.qos != QoS::AtMostOnce { - return self.queue_publish_message(topic, payload, &options).await; + return self + .queue_publish_message(topic, payload, &options) + .await + .map(StagedPublish::Queued); } - let options = self.resolve_effective_qos(options); - #[cfg(feature = "opentelemetry")] let options = { let mut opts = options; @@ -738,115 +920,96 @@ impl DirectClientInner { return Err(MqttError::NotConnected); } + if let Some(alias) = options.properties.topic_alias { + let session = self.session.read().await; + let aliases = session.topic_alias_out().read().await; + outbound::check_topic_alias(&aliases, &topic, alias)?; + } let (final_payload, properties) = self.encode_payload(payload, &options)?; - let needs_packet_id = options.qos != QoS::AtMostOnce; - - let mut publish = PublishPacket { + let mut publish = self.publish_policy().conform(PublishPacket { topic_name: topic, payload: final_payload, qos: options.qos, retain: options.retain, dup: false, - packet_id: needs_packet_id.then_some(SIZE_PROBE_PACKET_ID), + packet_id: (options.qos != QoS::AtMostOnce).then_some(SIZE_PROBE_PACKET_ID), properties, protocol_version: self.options.protocol_version.as_u8(), stream_id: None, - }; + })?; - let mut buf = bytes::BytesMut::new(); - publish.encode(&mut buf)?; - self.session - .read() - .await - .check_packet_size(buf.len()) - .await?; + self.check_publish_size(&publish).await?; - let packet_id = needs_packet_id.then(|| self.packet_id_generator.next()); - publish.packet_id = packet_id; + if publish.qos != QoS::AtMostOnce { + publish.packet_id = Some(self.allocate_packet_id().await?); + } - if let Some(pid) = packet_id { - let flow = self.session.read().await.flow_control().clone(); - flow.read().await.acquire_send_quota(pid).await?; + Ok(StagedPublish::Ready(publish)) + } + + pub(crate) async fn transmit_publish( + &self, + publish: PublishPacket, + ) -> Result> { + let qos = publish.qos; + let packet_id = publish.packet_id; + let flow = Arc::clone(self.session.read().await.flow_control()); + + if !self.is_connected() { + if let Some(pid) = packet_id { + Self::release_send_quota(&flow, pid).await; + } + return Err(MqttError::NotConnected); } - if options.qos != QoS::AtMostOnce { - if let Err(e) = self - .session - .write() - .await - .store_unacked_publish(publish.clone()) - .await - { - Self::release_outbound_quota(&self.session, packet_id).await; + if qos != QoS::AtMostOnce { + let stored = match self.with_aliased_topic(publish.clone()).await { + Ok(stored) => { + self.session + .read() + .await + .store_unacked_publish(stored) + .await + } + Err(e) => Err(e), + }; + if let Err(e) = stored { + if let Some(pid) = packet_id { + Self::release_send_quota(&flow, pid).await; + } return Err(e); } } - let rx = self.setup_publish_acknowledgment(options.qos, packet_id); + let ack = self.setup_publish_acknowledgment(qos, packet_id); if publish.payload.len() > 10000 { tracing::debug!( topic = %publish.topic_name, payload_len = publish.payload.len(), packet_id = ?packet_id, - qos = ?options.qos, + qos = ?qos, "Sending large PUBLISH packet" ); } - if let Err(e) = self.send_publish_packet(publish, options.qos).await { - Self::release_outbound_quota(&self.session, packet_id).await; - return Err(e); + let alias_mapping = publish + .topic_alias() + .filter(|_| !publish.topic_name.is_empty()) + .map(|alias| (alias, publish.topic_name.clone())); + self.send_publish_packet(publish).await?; + if let Some((alias, alias_topic)) = alias_mapping { + self.record_outbound_topic_alias(alias, &alias_topic).await; } - - if let Some(rx) = rx { - if let Err(e) = self - .wait_for_acknowledgment(rx, options.qos, packet_id) - .await - { - if !matches!(e, MqttError::Timeout) { - Self::release_outbound_quota(&self.session, packet_id).await; - } - return Err(e); - } - } - - Ok(match packet_id { - None => PublishResult::QoS0, - Some(id) => PublishResult::QoS1Or2 { packet_id: id }, - }) + Ok(ack) } - fn resolve_effective_qos(&self, options: PublishOptions) -> PublishOptions { - let effective_qos = if let Some(max_qos) = *self.server_max_qos.lock() { - let qos_value = match options.qos { - QoS::AtMostOnce => 0, - QoS::AtLeastOnce => 1, - QoS::ExactlyOnce => 2, - }; - if qos_value > max_qos { - tracing::warn!( - "Requested QoS {} exceeds server maximum {}, using QoS {}", - qos_value, - max_qos, - max_qos - ); - match max_qos { - 0 => QoS::AtMostOnce, - 1 => QoS::AtLeastOnce, - _ => QoS::ExactlyOnce, - } - } else { - options.qos - } - } else { - options.qos - }; - - PublishOptions { - qos: effective_qos, - ..options + async fn record_outbound_topic_alias(&self, alias: u16, topic: &str) { + let session = self.session.read().await; + let mut aliases = session.topic_alias_out().write().await; + if let Err(e) = aliases.register_alias(alias, topic) { + tracing::warn!(alias, topic, error = %e, "outbound Topic Alias not recorded"); } } @@ -864,19 +1027,16 @@ impl DirectClientInner { }; let mut properties: Properties = options.properties.clone().into(); - if let Some(ct) = codec_content_type { - use crate::protocol::v5::properties::{PropertyId, PropertyValue}; - let _ = properties.add(PropertyId::ContentType, PropertyValue::Utf8String(ct)); + if let Some(ct) = codec_content_type.filter(|_| properties.get_content_type().is_none()) { + properties.set_content_type(ct); } Ok((final_payload, properties)) } - async fn send_publish_packet(&self, publish: PublishPacket, qos: QoS) -> Result<()> { - #[cfg(not(feature = "transport-quic"))] - let _ = qos; - + async fn send_publish_packet(&self, publish: PublishPacket) -> Result<()> { #[cfg(feature = "transport-quic")] { + let qos = publish.qos; if qos == QoS::AtMostOnce && self.datagrams_available() { if let Some(max_size) = self.max_datagram_size() { let overhead = 5 + publish.topic_name.len(); @@ -911,12 +1071,12 @@ impl DirectClientInner { .await?; return Ok(()); } - #[allow(deprecated)] - StreamStrategy::DataPerTopic | StreamStrategy::DataPerSubscription => { + StreamStrategy::ControlOnly => {} + topic_strategy => { tracing::debug!( topic = %publish.topic_name, qos = ?qos, - strategy = ?manager.strategy(), + strategy = ?topic_strategy, "Using topic-specific QUIC stream for PUBLISH" ); manager @@ -927,7 +1087,6 @@ impl DirectClientInner { .await?; return Ok(()); } - StreamStrategy::ControlOnly => {} } } } @@ -942,41 +1101,36 @@ impl DirectClientInner { } #[cfg(feature = "transport-quic")] - #[allow(deprecated)] - async fn should_unsubscribe_on_data_flow(&self, packet: &UnsubscribePacket) -> bool { - if packet.filters.len() != 1 { - return false; - } - if let Some(manager) = &self.quic_stream_manager { - if !matches!( + fn topic_stream_manager(&self) -> Option<&Arc> { + self.quic_stream_manager.as_ref().filter(|manager| { + !matches!( manager.strategy(), - StreamStrategy::DataPerTopic | StreamStrategy::DataPerSubscription - ) { - return false; - } - manager - .get_flow_id_for_topic(&packet.filters[0]) - .await - .is_some() - } else { - false - } + StreamStrategy::ControlOnly | StreamStrategy::DataPerPublish + ) + }) + } + + #[cfg(feature = "transport-quic")] + async fn unsubscribe_data_flow_manager( + &self, + packet: &UnsubscribePacket, + ) -> Option<&Arc> { + let [filter] = packet.filters.as_slice() else { + return None; + }; + let manager = self.topic_stream_manager()?; + manager.get_flow_id_for_topic(filter).await.map(|_| manager) } #[cfg(feature = "transport-quic")] - #[allow(deprecated)] - fn should_subscribe_on_data_flow(&self, packet: &SubscribePacket) -> bool { + fn subscribe_data_flow_manager( + &self, + packet: &SubscribePacket, + ) -> Option<&Arc> { if packet.filters.len() != 1 { - return false; - } - if let Some(manager) = &self.quic_stream_manager { - matches!( - manager.strategy(), - StreamStrategy::DataPerTopic | StreamStrategy::DataPerSubscription - ) - } else { - false + return None; } + self.topic_stream_manager() } #[cfg(feature = "transport-quic")] @@ -1088,9 +1242,12 @@ impl DirectClientInner { return Err(MqttError::NotConnected); } + self.server_capabilities.check_subscribe(&packet)?; + self.check_packet_fits(&packet).await?; + let writer = self.writer.as_ref().ok_or(MqttError::NotConnected)?; - let packet_id = self.packet_id_generator.next(); + let packet_id = self.allocate_packet_id().await?; let mut packet = packet; packet.packet_id = packet_id; @@ -1106,8 +1263,7 @@ impl DirectClientInner { ); #[cfg(feature = "transport-quic")] - let sent_on_flow = if self.should_subscribe_on_data_flow(&packet) { - let manager = self.quic_stream_manager.as_ref().unwrap(); + let sent_on_flow = if let Some(manager) = self.subscribe_data_flow_manager(&packet) { let topic = packet.filters[0].filter.clone(); manager .send_on_topic_stream(topic, Packet::Subscribe(packet.clone())) @@ -1132,12 +1288,15 @@ impl DirectClientInner { for (filter, reason_code) in packet.filters.iter().zip(suback.reason_codes.iter()) { if let Some(subscription) = Self::create_subscription_from_filter(filter, *reason_code) { - self.session + let recorded = self + .session .write() .await .add_subscription(filter.filter.clone(), subscription) - .await - .ok(); + .await; + if let Err(e) = recorded { + tracing::warn!(filter = %filter.filter, error = %e, "subscription not recorded in session"); + } } } @@ -1162,9 +1321,11 @@ impl DirectClientInner { return Err(MqttError::NotConnected); } + self.check_unsubscribe(&packet).await?; + let writer = self.writer.as_ref().ok_or(MqttError::NotConnected)?; - let packet_id = self.packet_id_generator.next(); + let packet_id = self.allocate_packet_id().await?; let mut packet = packet; packet.packet_id = packet_id; @@ -1179,8 +1340,8 @@ impl DirectClientInner { } #[cfg(feature = "transport-quic")] - let sent_on_flow = if self.should_unsubscribe_on_data_flow(&packet).await { - let manager = self.quic_stream_manager.as_ref().unwrap(); + let sent_on_flow = if let Some(manager) = self.unsubscribe_data_flow_manager(&packet).await + { let topic = packet.filters[0].clone(); manager .send_on_topic_stream(topic, Packet::Unsubscribe(packet.clone())) @@ -1222,53 +1383,48 @@ impl DirectClientInner { } for filter in packet.filters { - let _ = self + let removed = self .session .write() .await .remove_subscription(&filter) .await; + if let Err(e) = removed { + tracing::warn!(filter = %filter, error = %e, "subscription not removed from session"); + } } Ok(()) } pub(crate) async fn build_connect_packet(&self) -> ConnectPacket { - use crate::protocol::v5::properties::{PropertyId, PropertyValue}; - let session = self.session.read().await; let mut properties = Properties::default(); if let Some(val) = self.options.properties.session_expiry_interval { - let _ = properties.add( - PropertyId::SessionExpiryInterval, - PropertyValue::FourByteInteger(val), - ); + properties.set_session_expiry_interval(val); } if let Some(val) = self.options.properties.receive_maximum { - let _ = properties.add( - PropertyId::ReceiveMaximum, - PropertyValue::TwoByteInteger(val), - ); + properties.set_receive_maximum(val); } if let Some(val) = self.options.properties.maximum_packet_size { - let _ = properties.add( - PropertyId::MaximumPacketSize, - PropertyValue::FourByteInteger(val), - ); + properties.set_maximum_packet_size(val); } if let Some(val) = self.options.properties.topic_alias_maximum { - let _ = properties.add( - PropertyId::TopicAliasMaximum, - PropertyValue::TwoByteInteger(val), - ); + properties.set_topic_alias_maximum(val); + } + if let Some(val) = self.options.properties.request_response_information { + properties.set_request_response_information(val); + } + if let Some(val) = self.options.properties.request_problem_information { + properties.set_request_problem_information(val); + } + for (key, value) in &self.options.properties.user_properties { + properties.add_user_property(key.clone(), value.clone()); } if let Some(ref method) = self.options.properties.authentication_method { - let _ = properties.add( - PropertyId::AuthenticationMethod, - PropertyValue::Utf8String(method.clone()), - ); + properties.set_authentication_method(method.clone()); let auth_data = if let Some(ref handler) = self.auth_handler { match handler.initial_response(method).await { @@ -1283,10 +1439,7 @@ impl DirectClientInner { }; if let Some(data) = auth_data { - let _ = properties.add( - PropertyId::AuthenticationData, - PropertyValue::BinaryData(bytes::Bytes::from(data)), - ); + properties.set_authentication_data(bytes::Bytes::from(data)); } } @@ -1349,6 +1502,14 @@ impl DirectClientInner { deferred_ack: self.options.deferred_ack, ack_callbacks: Arc::clone(&self.ack_callbacks), ack_dispatcher: Arc::clone(&self.ack_dispatcher), + topic_aliases: Arc::new(Mutex::new(crate::session::TopicAliasManager::new( + self.options.properties.topic_alias_maximum.unwrap_or(0), + ))), + request_problem_information: self + .options + .properties + .request_problem_information + .unwrap_or(true), }; let ctx_for_packet_reader = ctx.clone(); @@ -1442,12 +1603,6 @@ impl DirectClientInner { Ok(recovered) } - #[cfg(not(feature = "transport-quic"))] - #[allow(clippy::unused_async)] - pub(crate) async fn recover_flows(&self) -> crate::error::Result { - Ok(0) - } - #[cfg(feature = "transport-quic")] pub async fn discard_flow(&self, flow_id: FlowId) -> Result<()> { if !self.is_connected() { @@ -1592,7 +1747,7 @@ pub mod tests { let client = create_test_client(); let result = client - .publish( + .stage_publish( "test/topic".to_string(), b"test payload".to_vec(), PublishOptions::default(), @@ -1653,7 +1808,7 @@ pub mod tests { assert!(!client.is_connected()); let oversized = client - .publish( + .stage_publish( "test/flush".to_string(), vec![0u8; 4096], PublishOptions { @@ -1672,7 +1827,7 @@ pub mod tests { ); let within = client - .publish( + .stage_publish( "test/flush".to_string(), vec![0u8; 64], PublishOptions { @@ -1682,7 +1837,10 @@ pub mod tests { ) .await; assert!( - matches!(within, Ok(PublishResult::QoS1Or2 { .. })), + matches!( + within, + Ok(StagedPublish::Queued(PublishResult::QoS1Or2 { .. })) + ), "within-limit publish must queue: {within:?}" ); assert_eq!( diff --git a/crates/mqtt5/src/client/direct/outbound.rs b/crates/mqtt5/src/client/direct/outbound.rs new file mode 100644 index 00000000..c47313a0 --- /dev/null +++ b/crates/mqtt5/src/client/direct/outbound.rs @@ -0,0 +1,250 @@ +use crate::error::{MqttError, Result}; +use crate::packet::connack::ConnAckPacket; +use crate::packet::subscribe::SubscribePacket; +use crate::packet::unsubscribe::UnsubscribePacket; +use crate::protocol::v5::properties::{Properties, PropertyId, PropertyValue}; +use crate::session::TopicAliasManager; +use crate::types::PublishOptions; +use crate::validation::{ + parse_shared_subscription, validate_subscription_filter, validate_topic_name, +}; + +const MAX_SUBSCRIPTION_IDENTIFIER: u32 = 268_435_455; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct ServerCapabilities { + wildcard_subscriptions: bool, + shared_subscriptions: bool, + subscription_identifiers: bool, +} + +impl Default for ServerCapabilities { + fn default() -> Self { + Self { + wildcard_subscriptions: true, + shared_subscriptions: true, + subscription_identifiers: true, + } + } +} + +impl ServerCapabilities { + pub(crate) fn from_connack(connack: &ConnAckPacket) -> Self { + let properties = &connack.properties; + Self { + wildcard_subscriptions: is_available( + properties, + PropertyId::WildcardSubscriptionAvailable, + ), + shared_subscriptions: is_available(properties, PropertyId::SharedSubscriptionAvailable), + subscription_identifiers: is_available( + properties, + PropertyId::SubscriptionIdentifierAvailable, + ), + } + } + + pub(crate) fn check_subscribe(self, packet: &SubscribePacket) -> Result<()> { + let identifiers = packet.properties.subscription_identifiers(); + if let Some(invalid) = identifiers + .iter() + .find(|id| !(1..=MAX_SUBSCRIPTION_IDENTIFIER).contains(*id)) + { + return Err(MqttError::ProtocolError(format!( + "Subscription Identifier {invalid} is outside 1..={MAX_SUBSCRIPTION_IDENTIFIER}" + ))); + } + if !identifiers.is_empty() && !self.subscription_identifiers { + return Err(MqttError::SubscriptionIdentifiersNotSupported); + } + for filter in &packet.filters { + validate_subscription_filter(&filter.filter)?; + let (topic_filter, share_name) = parse_shared_subscription(&filter.filter); + if share_name.is_some() && !self.shared_subscriptions { + return Err(MqttError::SharedSubscriptionsNotSupported); + } + if topic_filter.contains(['+', '#']) && !self.wildcard_subscriptions { + return Err(MqttError::WildcardSubscriptionsNotSupported); + } + } + Ok(()) + } +} + +fn is_available(properties: &Properties, id: PropertyId) -> bool { + !matches!(properties.get(id), Some(PropertyValue::Byte(0))) +} + +pub(crate) fn check_unsubscribe(packet: &UnsubscribePacket) -> Result<()> { + packet + .filters + .iter() + .map(String::as_str) + .try_for_each(validate_subscription_filter) +} + +pub(crate) fn check_publish(topic: &str, options: &PublishOptions) -> Result<()> { + let properties = &options.properties; + match (topic.is_empty(), properties.topic_alias) { + (true, None) => { + return Err(MqttError::InvalidTopicName( + "zero-length Topic Name requires a Topic Alias".to_string(), + )); + } + (true, Some(_)) => {} + (false, _) => validate_topic_name(topic)?, + } + if properties.topic_alias == Some(0) { + return Err(MqttError::TopicAliasInvalid(0)); + } + if let Some(response_topic) = &properties.response_topic { + validate_topic_name(response_topic)?; + } + if !properties.subscription_identifiers.is_empty() { + return Err(MqttError::ProtocolError( + "a client PUBLISH must not contain a Subscription Identifier".to_string(), + )); + } + Ok(()) +} + +pub(crate) fn check_topic_alias( + aliases: &TopicAliasManager, + topic: &str, + alias: u16, +) -> Result<()> { + let known = if topic.is_empty() { + aliases.get_topic(alias).is_some() + } else { + (1..=aliases.topic_alias_maximum()).contains(&alias) + }; + if known { + Ok(()) + } else { + Err(MqttError::TopicAliasInvalid(alias)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::packet::subscribe::{SubscriptionOptions, TopicFilter}; + use crate::protocol::v5::reason_codes::ReasonCode; + use crate::types::PublishProperties; + + fn subscribe(filter: &str, identifier: Option) -> SubscribePacket { + let packet = SubscribePacket { + packet_id: 0, + filters: vec![TopicFilter { + filter: filter.to_string(), + options: SubscriptionOptions::default(), + }], + properties: Properties::default(), + protocol_version: 5, + }; + match identifier { + Some(id) => packet.with_subscription_identifier(id), + None => packet, + } + } + + fn capabilities(id: PropertyId) -> ServerCapabilities { + let mut connack = ConnAckPacket::new(false, ReasonCode::Success); + let _ = connack.properties.add(id, PropertyValue::Byte(0)); + ServerCapabilities::from_connack(&connack) + } + + fn publish_with(properties: PublishProperties) -> PublishOptions { + PublishOptions { + properties, + ..Default::default() + } + } + + #[test] + fn subscribe_respects_server_capabilities() { + let all = ServerCapabilities::default(); + assert!(all.check_subscribe(&subscribe("a/#", Some(1))).is_ok()); + assert!(all.check_subscribe(&subscribe("$share/g/a", None)).is_ok()); + + let no_wildcards = capabilities(PropertyId::WildcardSubscriptionAvailable); + assert!(matches!( + no_wildcards.check_subscribe(&subscribe("a/+", None)), + Err(MqttError::WildcardSubscriptionsNotSupported) + )); + assert!(no_wildcards + .check_subscribe(&subscribe("a/b", None)) + .is_ok()); + + let no_shared = capabilities(PropertyId::SharedSubscriptionAvailable); + assert!(matches!( + no_shared.check_subscribe(&subscribe("$share/g/a", None)), + Err(MqttError::SharedSubscriptionsNotSupported) + )); + + let no_ids = capabilities(PropertyId::SubscriptionIdentifierAvailable); + assert!(matches!( + no_ids.check_subscribe(&subscribe("a", Some(3))), + Err(MqttError::SubscriptionIdentifiersNotSupported) + )); + } + + #[test] + fn subscribe_rejects_identifier_zero_and_invalid_filters() { + let all = ServerCapabilities::default(); + assert!(all.check_subscribe(&subscribe("a", Some(0))).is_err()); + assert!(all.check_subscribe(&subscribe("a/#/b", None)).is_err()); + assert!(all.check_subscribe(&subscribe("$share/+/x", None)).is_err()); + } + + #[test] + fn publish_validation() { + let plain = PublishOptions::default(); + assert!(check_publish("a/b", &plain).is_ok()); + assert!(check_publish("a/+", &plain).is_err()); + assert!(check_publish("", &plain).is_err()); + assert!(check_publish( + "", + &publish_with(PublishProperties { + topic_alias: Some(1), + ..Default::default() + }) + ) + .is_ok()); + assert!(check_publish( + "a", + &publish_with(PublishProperties { + topic_alias: Some(0), + ..Default::default() + }) + ) + .is_err()); + assert!(check_publish( + "a", + &publish_with(PublishProperties { + response_topic: Some("r/#".to_string()), + ..Default::default() + }) + ) + .is_err()); + assert!(check_publish( + "a", + &publish_with(PublishProperties { + subscription_identifiers: vec![1], + ..Default::default() + }) + ) + .is_err()); + } + + #[test] + fn topic_alias_bounds_and_mapping() { + let mut aliases = TopicAliasManager::new(2); + assert!(check_topic_alias(&aliases, "a", 2).is_ok()); + assert!(check_topic_alias(&aliases, "a", 3).is_err()); + assert!(check_topic_alias(&aliases, "", 1).is_err()); + aliases.register_alias(1, "a").unwrap(); + assert!(check_topic_alias(&aliases, "", 1).is_ok()); + assert!(check_topic_alias(&TopicAliasManager::new(0), "a", 1).is_err()); + } +} diff --git a/crates/mqtt5/src/client/direct/reader.rs b/crates/mqtt5/src/client/direct/reader.rs index c534404d..f9b448d0 100644 --- a/crates/mqtt5/src/client/direct/reader.rs +++ b/crates/mqtt5/src/client/direct/reader.rs @@ -6,11 +6,13 @@ use crate::client::connection::DisconnectReason; use crate::codec::CodecRegistry; use crate::error::{MqttError, Result}; use crate::packet::auth::AuthPacket; +use crate::packet::disconnect::DisconnectPacket; use crate::packet::suback::SubAckPacket; use crate::packet::unsuback::UnsubAckPacket; use crate::packet::Packet; +use crate::protocol::v5::properties::{Properties, PropertyId}; use crate::protocol::v5::reason_codes::ReasonCode; -use crate::session::SessionState; +use crate::session::{SessionState, TopicAliasManager}; use crate::transport::PacketWriter; use parking_lot::Mutex; use std::collections::HashMap; @@ -34,6 +36,8 @@ use quinn::Connection; #[cfg(feature = "transport-quic")] use std::time::Duration as StdDuration; +const CLOSE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5); + #[derive(Clone)] pub(super) struct PacketReaderContext { pub(super) session: Arc>, @@ -53,6 +57,8 @@ pub(super) struct PacketReaderContext { pub(super) deferred_ack: bool, pub(super) ack_callbacks: Arc, pub(super) ack_dispatcher: Arc, + pub(super) topic_aliases: Arc>, + pub(super) request_problem_information: bool, } impl PacketReaderContext { @@ -73,6 +79,7 @@ impl PacketReaderContext { ) -> super::handlers::IncomingHandlers<'a> { super::handlers::IncomingHandlers { session: &self.session, + topic_aliases: &self.topic_aliases, callback_manager: &self.callback_manager, keepalive_state: &self.keepalive_state, codec_registry: self.codec_registry.as_ref(), @@ -108,99 +115,155 @@ fn disconnect_reason_for(error: &MqttError) -> DisconnectReason { } } +fn disconnect_code_for(error: &MqttError) -> Option { + match error { + MqttError::Io(_) + | MqttError::ConnectionError(_) + | MqttError::ConnectionClosedByPeer + | MqttError::ClientClosed + | MqttError::NotConnected + | MqttError::ServerDisconnect(_) + | MqttError::KeepAliveTimeout => None, + MqttError::MalformedPacket(_) + | MqttError::InvalidQoS(_) + | MqttError::InvalidPacketType(_) + | MqttError::InvalidReasonCode(_) + | MqttError::InvalidPropertyId(_) + | MqttError::InvalidTopicName(_) + | MqttError::StringTooLong(_) => Some(ReasonCode::MalformedPacket), + MqttError::ProtocolError(_) | MqttError::DuplicatePropertyId(_) => { + Some(ReasonCode::ProtocolError) + } + MqttError::PacketTooLarge { .. } => Some(ReasonCode::PacketTooLarge), + MqttError::ReceiveMaximumExceeded => Some(ReasonCode::ReceiveMaximumExceeded), + MqttError::TopicAliasInvalid(_) => Some(ReasonCode::TopicAliasInvalid), + MqttError::AuthenticationFailed | MqttError::NotAuthorized => { + Some(ReasonCode::NotAuthorized) + } + _ => Some(ReasonCode::UnspecifiedError), + } +} + +fn problem_information_properties(packet: &Packet) -> Option<&Properties> { + match packet { + Packet::PubAck(p) => Some(&p.properties), + Packet::PubRec(p) => Some(&p.properties), + Packet::PubRel(p) => Some(&p.properties), + Packet::PubComp(p) => Some(&p.properties), + Packet::SubAck(p) => Some(&p.properties), + Packet::UnsubAck(p) => Some(&p.properties), + Packet::Auth(p) => Some(&p.properties), + _ => None, + } +} + +fn check_problem_information(packet: &Packet, requested: bool) -> Result<()> { + let Some(properties) = problem_information_properties(packet).filter(|_| !requested) else { + return Ok(()); + }; + if properties.contains(PropertyId::ReasonString) + || properties.contains(PropertyId::UserProperty) + { + return Err(MqttError::ProtocolError(format!( + "{} carries a Reason String or User Property although Request Problem Information is 0 [MQTT-3.1.2-29]", + packet.packet_type_name() + ))); + } + Ok(()) +} + +pub(super) async fn close_connection( + writer: &Arc>, + error: &MqttError, +) { + let disconnect = disconnect_code_for(error) + .map(|reason_code| Packet::Disconnect(DisconnectPacket::new(reason_code))); + let closing = async { writer.lock().await.close(disconnect).await }; + match tokio::time::timeout(CLOSE_TIMEOUT, closing).await { + Ok(Ok(())) => {} + Ok(Err(e)) => tracing::debug!("Closing network connection: {e}"), + Err(_) => tracing::warn!("Timed out closing network connection"), + } +} + +async fn process_packet(packet: Packet, ctx: &PacketReaderContext) -> Result<()> { + tracing::trace!("Received packet: {:?}", packet); + check_problem_information(&packet, ctx.request_problem_information)?; + match &packet { + Packet::SubAck(suback) => { + if let Some(tx) = ctx.suback_channels.lock().remove(&suback.packet_id) { + let _ = tx.send(suback.clone()); + return Ok(()); + } + } + Packet::UnsubAck(unsuback) => { + if let Some(tx) = ctx.unsuback_channels.lock().remove(&unsuback.packet_id) { + let _ = tx.send(unsuback.clone()); + return Ok(()); + } + } + Packet::PubAck(puback) => { + super::DirectClientInner::release_outbound_quota(&ctx.session, Some(puback.packet_id)) + .await; + if let Some(tx) = ctx.puback_channels.lock().remove(&puback.packet_id) { + let _ = tx.send(puback.reason_code); + return Ok(()); + } + } + Packet::PubRec(pubrec) if pubrec.reason_code.is_error() => { + tracing::debug!( + packet_id = pubrec.packet_id, + reason_code = ?pubrec.reason_code, + "QoS 2 PUBREC rejected" + ); + if let Some(tx) = ctx.pubcomp_channels.lock().remove(&pubrec.packet_id) { + let _ = tx.send(pubrec.reason_code); + } + ctx.session + .write() + .await + .remove_unacked_publish(pubrec.packet_id) + .await; + super::DirectClientInner::release_outbound_quota(&ctx.session, Some(pubrec.packet_id)) + .await; + return Ok(()); + } + Packet::PubComp(pubcomp) => { + super::DirectClientInner::release_outbound_quota(&ctx.session, Some(pubcomp.packet_id)) + .await; + if let Some(tx) = ctx.pubcomp_channels.lock().remove(&pubcomp.packet_id) { + let _ = tx.send(pubcomp.reason_code); + } + return Ok(()); + } + Packet::Auth(auth) => return handle_auth_packet(auth.clone(), ctx).await, + _ => {} + } + + let ack_delivery = ctx.ack_delivery(); + let handlers = ctx.incoming_handlers(ack_delivery.as_ref()); + handle_incoming_packet_with_writer(packet, &ctx.writer, None, &handlers).await +} + pub(super) async fn packet_reader_task_with_responses( mut reader: UnifiedReader, ctx: PacketReaderContext, ) { tracing::debug!("Packet reader task started and ready to process incoming packets"); - let disconnect_reason = loop { - let packet = reader.read_packet().await; - - match packet { - Ok(packet) => { - tracing::trace!("Received packet: {:?}", packet); - match &packet { - Packet::SubAck(suback) => { - if let Some(tx) = ctx.suback_channels.lock().remove(&suback.packet_id) { - let _ = tx.send(suback.clone()); - continue; - } - } - Packet::UnsubAck(unsuback) => { - if let Some(tx) = ctx.unsuback_channels.lock().remove(&unsuback.packet_id) { - let _ = tx.send(unsuback.clone()); - continue; - } - } - Packet::PubAck(puback) => { - super::DirectClientInner::release_outbound_quota( - &ctx.session, - Some(puback.packet_id), - ) - .await; - if let Some(tx) = ctx.puback_channels.lock().remove(&puback.packet_id) { - let _ = tx.send(puback.reason_code); - continue; - } - } - Packet::PubRec(pubrec) if pubrec.reason_code.is_error() => { - tracing::debug!( - packet_id = pubrec.packet_id, - reason_code = ?pubrec.reason_code, - "QoS 2 PUBREC rejected" - ); - if let Some(tx) = ctx.pubcomp_channels.lock().remove(&pubrec.packet_id) { - let _ = tx.send(pubrec.reason_code); - } - ctx.session - .write() - .await - .remove_unacked_publish(pubrec.packet_id) - .await; - super::DirectClientInner::release_outbound_quota( - &ctx.session, - Some(pubrec.packet_id), - ) - .await; - continue; - } - Packet::PubComp(pubcomp) => { - super::DirectClientInner::release_outbound_quota( - &ctx.session, - Some(pubcomp.packet_id), - ) - .await; - if let Some(tx) = ctx.pubcomp_channels.lock().remove(&pubcomp.packet_id) { - let _ = tx.send(pubcomp.reason_code); - } - } - Packet::Auth(ref auth) => { - if let Err(e) = handle_auth_packet(auth.clone(), &ctx).await { - tracing::error!("Error handling AUTH packet: {e}"); - break disconnect_reason_for(&e); - } - continue; - } - _ => {} - } - - let ack_delivery = ctx.ack_delivery(); - let handlers = ctx.incoming_handlers(ack_delivery.as_ref()); - if let Err(e) = - handle_incoming_packet_with_writer(packet, &ctx.writer, None, &handlers).await - { - tracing::error!("Error handling packet: {e}"); - break disconnect_reason_for(&e); - } - } - Err(e) => { - tracing::error!("Error reading packet: {e}"); - break DisconnectReason::NetworkError(e.to_string()); - } + let failure = loop { + let outcome = match reader.read_packet().await { + Ok(packet) => process_packet(packet, &ctx).await, + Err(e) => Err(e), + }; + if let Err(e) = outcome { + break e; } }; - ctx.lifecycle.end(disconnect_reason).await; + tracing::error!("Packet reader stopping: {failure}"); + close_connection(&ctx.writer, &failure).await; + drop(reader); + ctx.lifecycle.end(disconnect_reason_for(&failure)).await; ctx.clear_pending_if_current(); } @@ -212,6 +275,11 @@ async fn handle_auth_packet(auth: AuthPacket, ctx: &PacketReaderContext) -> Resu match auth.reason_code { ReasonCode::ContinueAuthentication => { + let method = ctx.auth_method.clone().ok_or_else(|| { + MqttError::ProtocolError( + "AUTH received but CONNECT carried no Authentication Method".to_string(), + ) + })?; let handler = ctx .auth_handler .as_ref() @@ -224,7 +292,6 @@ async fn handle_auth_packet(auth: AuthPacket, ctx: &PacketReaderContext) -> Resu match response { AuthResponse::Continue(data) => { - let method = ctx.auth_method.clone().unwrap_or_default(); let auth_packet = AuthPacket::continue_authentication(method, Some(data))?; ctx.writer .lock() @@ -549,6 +616,7 @@ async fn quic_stream_reader_task( { tracing::error!(flow_id = ?flow_id, "Error handling packet from server stream: {e}"); if let MqttError::ServerDisconnect(reason_code) = e { + close_connection(&ctx.writer, &e).await; ctx.lifecycle .end(DisconnectReason::ServerDisconnect(reason_code)) .await; @@ -603,6 +671,7 @@ async fn quic_uni_stream_reader_task(mut recv: quinn::RecvStream, ctx: PacketRea { tracing::error!(flow_id = ?flow_id, "Error handling packet from uni stream: {e}"); if let MqttError::ServerDisconnect(reason_code) = e { + close_connection(&ctx.writer, &e).await; ctx.lifecycle .end(DisconnectReason::ServerDisconnect(reason_code)) .await; diff --git a/crates/mqtt5/src/client/direct/replay.rs b/crates/mqtt5/src/client/direct/replay.rs new file mode 100644 index 00000000..320050c0 --- /dev/null +++ b/crates/mqtt5/src/client/direct/replay.rs @@ -0,0 +1,166 @@ +use crate::error::{MqttError, Result}; +use crate::packet::publish::PublishPacket; +use crate::packet::pubrel::PubRelPacket; +use crate::packet::{MqttPacket, Packet}; +use crate::session::flow_control::FlowControlManager; +use crate::session::state::OutboundReplay; +use crate::session::SessionState; +use crate::transport::PacketWriter; +use crate::QoS; +use parking_lot::Mutex; +use std::collections::VecDeque; +use std::sync::{Arc, Weak}; +use tokio::sync::{RwLock, Semaphore}; + +use super::unified::UnifiedWriter; + +#[derive(Debug, Clone, Copy)] +pub(crate) struct PublishPolicy { + pub(crate) maximum_qos: Option, + pub(crate) retain_available: bool, +} + +impl PublishPolicy { + pub(crate) fn conform(self, mut publish: PublishPacket) -> Result { + if publish.retain && !self.retain_available { + return Err(MqttError::RetainNotSupported); + } + let requested = publish.qos as u8; + if let Some(maximum) = self.maximum_qos.filter(|maximum| requested > *maximum) { + tracing::warn!( + "Requested QoS {requested} exceeds server maximum {maximum}, using QoS {maximum}" + ); + publish.qos = match maximum { + 0 => QoS::AtMostOnce, + 1 => QoS::AtLeastOnce, + _ => QoS::ExactlyOnce, + }; + if publish.qos == QoS::AtMostOnce { + publish.packet_id = None; + } + } + Ok(publish) + } +} + +pub(super) struct SessionReplay { + pub(super) items: Vec, + pub(super) slots: Arc, + pub(super) session: Arc>, + pub(super) writer: Weak>, + pub(super) queued: Arc>>, + pub(super) policy: PublishPolicy, +} + +impl SessionReplay { + pub(super) async fn run(self) { + let flow = Arc::clone(self.session.read().await.flow_control()); + if self.replay_session_state(&flow).await && self.flush_offline_queue(&flow).await { + flow.read().await.finish_replay(&self.slots).await; + tracing::debug!("Session replay complete"); + } + } + + async fn replay_session_state(&self, flow: &Arc>) -> bool { + for item in &self.items { + let packet = match item { + OutboundReplay::PubRel(packet_id) => Packet::PubRel(PubRelPacket::new(*packet_id)), + OutboundReplay::Publish(publish) => { + if !self.take_slot(flow, publish.packet_id).await { + return false; + } + let mut resend = without_topic_alias(publish.clone()); + resend.dup = true; + Packet::Publish(resend) + } + }; + if !self.write(packet).await { + return false; + } + } + true + } + + async fn flush_offline_queue(&self, flow: &Arc>) -> bool { + loop { + let Some(queued) = self.queued.lock().front().cloned() else { + return true; + }; + let publish = match self.conform_queued(queued).await { + Ok(publish) => publish, + Err(e) => { + tracing::warn!("Dropping queued message: {e}"); + self.queued.lock().pop_front(); + continue; + } + }; + if !self.take_slot(flow, publish.packet_id).await { + return false; + } + self.queued.lock().pop_front(); + if publish.qos != QoS::AtMostOnce { + if let Err(e) = self + .session + .read() + .await + .store_unacked_publish(publish.clone()) + .await + { + tracing::warn!("Dropping queued message: {e}"); + continue; + } + } + if !self.write(Packet::Publish(publish)).await { + return false; + } + } + } + + async fn conform_queued(&self, queued: PublishPacket) -> Result { + let publish = without_topic_alias(self.policy.conform(queued)?); + let mut buf = bytes::BytesMut::new(); + publish.encode(&mut buf)?; + self.session + .read() + .await + .check_packet_size(buf.len()) + .await?; + Ok(publish) + } + + async fn take_slot( + &self, + flow: &Arc>, + packet_id: Option, + ) -> bool { + let Some(packet_id) = packet_id else { + return true; + }; + match self.slots.acquire().await { + Ok(permit) => { + permit.forget(); + flow.read() + .await + .claim_send_quota(&self.slots, packet_id) + .await + } + Err(_) => false, + } + } + + async fn write(&self, packet: Packet) -> bool { + let Some(writer) = self.writer.upgrade() else { + return false; + }; + let written = writer.lock().await.write_packet(packet).await; + if let Err(e) = &written { + tracing::debug!("Session replay stopped: {e}"); + } + written.is_ok() + } +} + +fn without_topic_alias(mut publish: PublishPacket) -> PublishPacket { + publish.properties.remove_topic_alias(); + publish +} diff --git a/crates/mqtt5/src/client/direct/unified.rs b/crates/mqtt5/src/client/direct/unified.rs index 8f9a944c..d88fa818 100644 --- a/crates/mqtt5/src/client/direct/unified.rs +++ b/crates/mqtt5/src/client/direct/unified.rs @@ -1,14 +1,19 @@ //! Unified reader and writer types for all transport types -use crate::error::Result; +use crate::error::{MqttError, Result}; use crate::packet::Packet; +use crate::transport::packet_io::read_packet_from_stream; use crate::transport::tls::{TlsReadHalf, TlsWriteHalf}; -use crate::transport::{PacketReader, PacketWriter}; +use crate::transport::PacketWriter; +use bytes::BytesMut; +use tokio::io::AsyncWriteExt; use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf}; #[cfg(feature = "transport-websocket")] use crate::transport::websocket::{WebSocketReadHandle, WebSocketWriteHandle}; #[cfg(feature = "transport-quic")] +use crate::transport::PacketReader; +#[cfg(feature = "transport-quic")] use quinn::{RecvStream, SendStream}; enum UnifiedReaderInner { @@ -23,6 +28,8 @@ enum UnifiedReaderInner { pub struct UnifiedReader { inner: UnifiedReaderInner, protocol_version: u8, + read_buffer: BytesMut, + max_packet_size: usize, } impl UnifiedReader { @@ -30,6 +37,8 @@ impl UnifiedReader { Self { inner: UnifiedReaderInner::Tcp(reader), protocol_version, + read_buffer: BytesMut::new(), + max_packet_size: usize::MAX, } } @@ -37,6 +46,8 @@ impl UnifiedReader { Self { inner: UnifiedReaderInner::Tls(reader), protocol_version, + read_buffer: BytesMut::new(), + max_packet_size: usize::MAX, } } @@ -45,6 +56,8 @@ impl UnifiedReader { Self { inner: UnifiedReaderInner::WebSocket(reader), protocol_version, + read_buffer: BytesMut::new(), + max_packet_size: usize::MAX, } } @@ -53,16 +66,44 @@ impl UnifiedReader { Self { inner: UnifiedReaderInner::Quic(reader), protocol_version, + read_buffer: BytesMut::new(), + max_packet_size: usize::MAX, } } + #[must_use] + pub fn with_maximum_packet_size(mut self, maximum_packet_size: Option) -> Self { + self.max_packet_size = maximum_packet_size + .and_then(|size| usize::try_from(size).ok()) + .unwrap_or(usize::MAX); + self + } + pub async fn read_packet(&mut self) -> Result { match &mut self.inner { - UnifiedReaderInner::Tcp(reader) => reader.read_packet(self.protocol_version).await, - UnifiedReaderInner::Tls(reader) => reader.read_packet(self.protocol_version).await, + UnifiedReaderInner::Tcp(reader) => { + read_packet_from_stream( + reader, + self.protocol_version, + &mut self.read_buffer, + self.max_packet_size, + ) + .await + } + UnifiedReaderInner::Tls(reader) => { + read_packet_from_stream( + reader, + self.protocol_version, + &mut self.read_buffer, + self.max_packet_size, + ) + .await + } #[cfg(feature = "transport-websocket")] UnifiedReaderInner::WebSocket(reader) => { - reader.read_packet(self.protocol_version).await + reader + .read_packet_limited(self.protocol_version, self.max_packet_size) + .await } #[cfg(feature = "transport-quic")] UnifiedReaderInner::Quic(reader) => reader.read_packet(self.protocol_version).await, @@ -77,6 +118,33 @@ pub enum UnifiedWriter { WebSocket(WebSocketWriteHandle), #[cfg(feature = "transport-quic")] Quic(SendStream), + Closed, +} + +impl UnifiedWriter { + pub async fn close(&mut self, final_packet: Option) -> Result<()> { + let mut current = std::mem::replace(self, Self::Closed); + let written = match final_packet { + Some(packet) => current.write_packet(packet).await, + None => Ok(()), + }; + let shutdown = current.shutdown().await; + written.and(shutdown) + } + + async fn shutdown(&mut self) -> Result<()> { + match self { + Self::Tcp(writer) => Ok(writer.shutdown().await?), + Self::Tls(writer) => Ok(writer.shutdown().await?), + #[cfg(feature = "transport-websocket")] + Self::WebSocket(writer) => writer.close().await, + #[cfg(feature = "transport-quic")] + Self::Quic(writer) => writer + .finish() + .map_err(|e| MqttError::ConnectionError(format!("QUIC stream finish: {e}"))), + Self::Closed => Ok(()), + } + } } impl PacketWriter for UnifiedWriter { @@ -88,6 +156,7 @@ impl PacketWriter for UnifiedWriter { Self::WebSocket(writer) => writer.write_packet(packet).await, #[cfg(feature = "transport-quic")] Self::Quic(writer) => writer.write_packet(packet).await, + Self::Closed => Err(MqttError::NotConnected), } } } diff --git a/crates/mqtt5/src/client/inner.rs b/crates/mqtt5/src/client/inner.rs index 1a25dea2..552d43af 100644 --- a/crates/mqtt5/src/client/inner.rs +++ b/crates/mqtt5/src/client/inner.rs @@ -30,14 +30,12 @@ impl MqttClient { session_present: bool, keep_alive: std::time::Duration, ) { - if !session_present { - self.reset_inbound_state_for_lost_session().await; - } self.trigger_connection_event(ConnectionEvent::Connected { session_present, keep_alive, }) .await; + #[cfg(feature = "transport-quic")] self.recover_quic_flows().await; self.restore_subscriptions_after_connect(stored_subs, session_present) .await; @@ -46,30 +44,6 @@ impl MqttClient { } } - /// Clears stale inbound `QoS` 2 de-duplication state after the broker reports no session. - /// - /// Without a resumed session the broker's packet-ID space starts fresh, so any retained - /// `inbound_delivered` guard would suppress a genuinely new PUBLISH that reuses the ID. - /// With deferred ack any outstanding [`AckToken`](crate::AckToken) is also stale — the - /// message it referred to is gone — so this is surfaced as a warning. - async fn reset_inbound_state_for_lost_session(&self) { - let inner = self.inner.read().await; - let cleared = inner.session.read().await.clear_all_inbound_state().await; - if cleared { - if inner.options.deferred_ack { - tracing::warn!( - "Reconnected with session_present=0; cleared stale inbound QoS 2 \ - de-duplication state. Any outstanding AckTokens are now stale because the \ - broker no longer holds the session that delivered their messages." - ); - } else { - tracing::debug!( - "Reconnected with session_present=0; cleared stale inbound QoS 2 de-duplication state" - ); - } - } - } - #[cfg(any(not(feature = "transport-websocket"), not(feature = "transport-quic")))] fn unsupported_transport_feature(transport: &str, feature: &str) -> MqttError { MqttError::Configuration(format!( @@ -94,9 +68,9 @@ impl MqttClient { let (host, port) = Self::split_host_port(rest, 8883)?; Ok((ClientTransportType::Tls, host, port)) } else if let Some(rest) = address.strip_prefix("ws://") { - let (host, port) = Self::split_host_port(rest, 80)?; #[cfg(feature = "transport-websocket")] { + let (host, port) = Self::split_host_port(rest, 80)?; Ok(( ClientTransportType::WebSocket(address.to_string()), host, @@ -105,16 +79,15 @@ impl MqttClient { } #[cfg(not(feature = "transport-websocket"))] { - let _ = (host, port); - Err(Self::unsupported_transport_feature( + Self::split_host_port(rest, 80).and(Err(Self::unsupported_transport_feature( "WebSocket", "transport-websocket", - )) + ))) } } else if let Some(rest) = address.strip_prefix("wss://") { - let (host, port) = Self::split_host_port(rest, 443)?; #[cfg(feature = "transport-websocket")] { + let (host, port) = Self::split_host_port(rest, 443)?; Ok(( ClientTransportType::WebSocketSecure(address.to_string()), host, @@ -123,11 +96,10 @@ impl MqttClient { } #[cfg(not(feature = "transport-websocket"))] { - let _ = (host, port); - Err(Self::unsupported_transport_feature( + Self::split_host_port(rest, 443).and(Err(Self::unsupported_transport_feature( "WebSocket", "transport-websocket", - )) + ))) } } else if let Some(rest) = address.strip_prefix("tcp://") { let (host, port) = Self::split_host_port(rest, 1883)?; @@ -136,32 +108,30 @@ impl MqttClient { let (host, port) = Self::split_host_port(rest, 8883)?; Ok((ClientTransportType::Tls, host, port)) } else if let Some(rest) = address.strip_prefix("quic://") { - let (host, port) = Self::split_host_port(rest, 14567)?; #[cfg(feature = "transport-quic")] { + let (host, port) = Self::split_host_port(rest, 14567)?; Ok((ClientTransportType::Quic, host, port)) } #[cfg(not(feature = "transport-quic"))] { - let _ = (host, port); - Err(Self::unsupported_transport_feature( + Self::split_host_port(rest, 14567).and(Err(Self::unsupported_transport_feature( "QUIC", "transport-quic", - )) + ))) } } else if let Some(rest) = address.strip_prefix("quics://") { - let (host, port) = Self::split_host_port(rest, 14567)?; #[cfg(feature = "transport-quic")] { + let (host, port) = Self::split_host_port(rest, 14567)?; Ok((ClientTransportType::QuicSecure, host, port)) } #[cfg(not(feature = "transport-quic"))] { - let _ = (host, port); - Err(Self::unsupported_transport_feature( + Self::split_host_port(rest, 14567).and(Err(Self::unsupported_transport_feature( "QUIC", "transport-quic", - )) + ))) } } else { let (host, port) = Self::split_host_port(address, 1883)?; diff --git a/crates/mqtt5/src/client/mod.rs b/crates/mqtt5/src/client/mod.rs index 4d2218d6..73714610 100644 --- a/crates/mqtt5/src/client/mod.rs +++ b/crates/mqtt5/src/client/mod.rs @@ -65,7 +65,8 @@ pub(crate) async fn fire_connection_event( pub use self::direct::AckToken; use self::direct::AutomaticReconnectLifecycle; #[cfg(not(target_arch = "wasm32"))] -use self::direct::DirectClientInner; +use self::direct::{DirectClientInner, StagedPublish}; +use crate::session::flow_control::FlowControlManager; /// Thread-safe MQTT v5.0 client /// @@ -419,8 +420,18 @@ impl MqttClient { "Publishing MQTT message" ); - let inner = self.inner.read().await; - match inner.publish(topic_str.clone(), payload_vec, options).await { + let staged = self + .inner + .read() + .await + .stage_publish(topic_str.clone(), payload_vec, options) + .await; + let outcome = match staged { + Ok(StagedPublish::Queued(result)) => Ok(result), + Ok(StagedPublish::Ready(publish)) => self.send_staged_publish(publish).await, + Err(e) => Err(e), + }; + match outcome { Ok(result) => { match &result { PublishResult::QoS0 => { @@ -439,6 +450,23 @@ impl MqttClient { } } + async fn send_staged_publish(&self, publish: PublishPacket) -> Result { + let packet_id = publish.packet_id; + if let Some(pid) = packet_id { + let flow = Arc::clone(self.inner.read().await.session.read().await.flow_control()); + FlowControlManager::acquire_shared_send_quota(&flow, pid).await?; + } + let ack = self.inner.read().await.transmit_publish(publish).await?; + if let Some(ack) = ack { + ack.wait().await?; + } + Ok( + packet_id.map_or(PublishResult::QoS0, |packet_id| PublishResult::QoS1Or2 { + packet_id, + }), + ) + } + /// Subscribes to a topic with a callback /// /// # Errors @@ -705,19 +733,20 @@ impl MqttClient { ); let inner = self.inner.read().await; - let _ = inner.callback_manager.unregister(&topic_filter); - let _ = inner.ack_callbacks.unregister(&topic_filter); - inner - .stored_ack_subscriptions - .lock() - .retain(|(topic, _, _, _)| topic != &topic_filter); - let packet = UnsubscribePacket { packet_id: 0, filters: vec![topic_filter.clone()], properties: Properties::default(), protocol_version: inner.options.protocol_version.as_u8(), }; + inner.check_unsubscribe(&packet).await?; + + let _ = inner.callback_manager.unregister(&topic_filter); + let _ = inner.ack_callbacks.unregister(&topic_filter); + inner + .stored_ack_subscriptions + .lock() + .retain(|(topic, _, _, _)| topic != &topic_filter); match inner.unsubscribe(packet).await { Ok(()) => { @@ -934,18 +963,17 @@ impl MqttClient { } } -#[allow(clippy::manual_async_fn)] impl MqttClientTrait for MqttClient { fn is_connected(&self) -> impl Future + Send + '_ { - async move { self.is_connected().await } + self.is_connected() } fn client_id(&self) -> impl Future + Send + '_ { - async move { self.client_id().await } + self.client_id() } fn connect<'a>(&'a self, address: &'a str) -> impl Future> + Send + 'a { - async move { self.connect(address).await } + self.connect(address) } fn connect_with_options<'a>( @@ -953,11 +981,11 @@ impl MqttClientTrait for MqttClient { address: &'a str, options: ConnectOptions, ) -> impl Future> + Send + 'a { - async move { Box::pin(self.connect_with_options(address, options)).await } + Box::pin(self.connect_with_options(address, options)) } fn disconnect(&self) -> impl Future> + Send + '_ { - async move { self.disconnect().await } + self.disconnect() } fn publish<'a>( @@ -965,7 +993,7 @@ impl MqttClientTrait for MqttClient { topic: impl Into + Send + 'a, payload: impl Into> + Send + 'a, ) -> impl Future> + Send + 'a { - async move { self.publish(topic, payload).await } + self.publish(topic, payload) } fn publish_qos<'a>( @@ -974,7 +1002,7 @@ impl MqttClientTrait for MqttClient { payload: impl Into> + Send + 'a, qos: QoS, ) -> impl Future> + Send + 'a { - async move { self.publish_qos(topic, payload, qos).await } + self.publish_qos(topic, payload, qos) } fn publish_with_options<'a>( @@ -983,7 +1011,7 @@ impl MqttClientTrait for MqttClient { payload: impl Into> + Send + 'a, options: PublishOptions, ) -> impl Future> + Send + 'a { - async move { self.publish_with_options(topic, payload, options).await } + self.publish_with_options(topic, payload, options) } fn subscribe<'a, F>( @@ -994,7 +1022,7 @@ impl MqttClientTrait for MqttClient { where F: Fn(crate::types::Message) + Send + Sync + 'static, { - async move { self.subscribe(topic_filter, callback).await } + self.subscribe(topic_filter, callback) } fn subscribe_with_options<'a, F>( @@ -1006,17 +1034,14 @@ impl MqttClientTrait for MqttClient { where F: Fn(crate::types::Message) + Send + Sync + 'static, { - async move { - self.subscribe_with_options(topic_filter, options, callback) - .await - } + self.subscribe_with_options(topic_filter, options, callback) } fn unsubscribe<'a>( &'a self, topic_filter: impl Into + Send + 'a, ) -> impl Future> + Send + 'a { - async move { self.unsubscribe(topic_filter).await } + self.unsubscribe(topic_filter) } fn subscribe_many<'a, F>( @@ -1027,14 +1052,14 @@ impl MqttClientTrait for MqttClient { where F: Fn(crate::types::Message) + Send + Sync + 'static + Clone, { - async move { self.subscribe_many(topics, callback).await } + self.subscribe_many(topics, callback) } fn unsubscribe_many<'a>( &'a self, topics: Vec<&'a str>, ) -> impl Future)>>> + Send + 'a { - async move { self.unsubscribe_many(topics).await } + self.unsubscribe_many(topics) } fn publish_retain<'a>( @@ -1042,7 +1067,7 @@ impl MqttClientTrait for MqttClient { topic: impl Into + Send + 'a, payload: impl Into> + Send + 'a, ) -> impl Future> + Send + 'a { - async move { self.publish_retain(topic, payload).await } + self.publish_retain(topic, payload) } fn publish_qos0<'a>( @@ -1050,7 +1075,7 @@ impl MqttClientTrait for MqttClient { topic: impl Into + Send + 'a, payload: impl Into> + Send + 'a, ) -> impl Future> + Send + 'a { - async move { self.publish_qos0(topic, payload).await } + self.publish_qos0(topic, payload) } fn publish_qos1<'a>( @@ -1058,7 +1083,7 @@ impl MqttClientTrait for MqttClient { topic: impl Into + Send + 'a, payload: impl Into> + Send + 'a, ) -> impl Future> + Send + 'a { - async move { self.publish_qos1(topic, payload).await } + self.publish_qos1(topic, payload) } fn publish_qos2<'a>( @@ -1066,15 +1091,15 @@ impl MqttClientTrait for MqttClient { topic: impl Into + Send + 'a, payload: impl Into> + Send + 'a, ) -> impl Future> + Send + 'a { - async move { self.publish_qos2(topic, payload).await } + self.publish_qos2(topic, payload) } fn is_queue_on_disconnect(&self) -> impl Future + Send + '_ { - async move { self.is_queue_on_disconnect().await } + self.is_queue_on_disconnect() } fn set_queue_on_disconnect(&self, enabled: bool) -> impl Future + Send + '_ { - async move { self.set_queue_on_disconnect(enabled).await } + self.set_queue_on_disconnect(enabled) } } diff --git a/crates/mqtt5/src/client/state.rs b/crates/mqtt5/src/client/state.rs index ba8ebb16..ebf9e6bb 100644 --- a/crates/mqtt5/src/client/state.rs +++ b/crates/mqtt5/src/client/state.rs @@ -36,6 +36,7 @@ impl MqttClient { inner.reconnect_attempt = 0; } + #[cfg(feature = "transport-quic")] pub(crate) async fn recover_quic_flows(&self) { let inner = self.inner.read().await; match inner.recover_flows().await { @@ -120,7 +121,7 @@ impl MqttClient { properties.set_subscription_identifier(id); } let packet = SubscribePacket { - packet_id: inner.packet_id_generator.next(), + packet_id: 0, filters: vec![crate::packet::subscribe::TopicFilter { filter: topic.to_string(), options, @@ -279,48 +280,19 @@ impl MqttClient { match reconnection_result { Ok(_) => { tracing::info!("Reconnected successfully after {} attempts", attempt); - - self.send_queued_messages().await; - return Ok(()); } Err(e) => { tracing::warn!("Reconnection attempt {} failed: {}", attempt, e); - #[allow(clippy::cast_possible_truncation)] - { - delay = std::cmp::min( - Duration::from_secs_f64(delay.as_secs_f64() * config.backoff_factor()), - config.max_delay, - ); - } + delay = + Duration::try_from_secs_f64(delay.as_secs_f64() * config.backoff_factor()) + .map_or(config.max_delay, |next| next.min(config.max_delay)); } } } } - pub(crate) async fn send_queued_messages(&self) { - let messages = { - let inner = self.inner.read().await; - let mut queued = inner.queued_messages.lock(); - std::mem::take(&mut *queued) - }; - - for mut msg in messages { - msg.dup = true; - - let size_check = self.inner.read().await.check_publish_size(&msg).await; - if let Err(e) = size_check { - tracing::warn!("Dropping queued message exceeding negotiated packet size: {e}"); - continue; - } - - if let Err(e) = self.publish_packet(msg).await { - tracing::warn!("Failed to send queued message: {e}"); - } - } - } - /// Internal method to resubscribe with stored options and callback /// /// # Errors @@ -343,7 +315,7 @@ impl MqttClient { properties.set_subscription_identifier(id); } let packet = SubscribePacket { - packet_id: inner.packet_id_generator.next(), + packet_id: 0, filters: vec![crate::packet::subscribe::TopicFilter { filter: topic.to_string(), options, @@ -356,19 +328,4 @@ impl MqttClient { .await?; Ok(()) } - - /// Internal method to publish a packet - /// - /// # Errors - /// - /// Returns an error if publish fails - pub(crate) async fn publish_packet( - &self, - packet: crate::packet::publish::PublishPacket, - ) -> Result<()> { - let mut inner = self.inner.write().await; - inner - .send_packet(crate::packet::Packet::Publish(packet)) - .await - } } diff --git a/crates/mqtt5/src/lib.rs b/crates/mqtt5/src/lib.rs index 09db8dbd..c7f7a90b 100644 --- a/crates/mqtt5/src/lib.rs +++ b/crates/mqtt5/src/lib.rs @@ -49,6 +49,7 @@ //! // Configure connection options //! let options = ConnectOptions::new("weather-station") //! .with_clean_start(false) // Resume previous session +//! .with_resume_existing_session(true) //! .with_keep_alive(Duration::from_secs(30)) //! .with_automatic_reconnect(true) //! .with_reconnect_delay(Duration::from_secs(5), Duration::from_secs(60)); @@ -202,6 +203,7 @@ //! let options = ConnectOptions::new("worker") //! .with_deferred_ack(true) //! .with_clean_start(false) +//! .with_resume_existing_session(true) //! .with_session_expiry_interval(3600) //! .with_receive_maximum(16); //! diff --git a/crates/mqtt5/src/session.rs b/crates/mqtt5/src/session.rs index 15a1a18a..5a98a999 100644 --- a/crates/mqtt5/src/session.rs +++ b/crates/mqtt5/src/session.rs @@ -3,7 +3,6 @@ pub mod limits; pub mod queue; #[cfg(not(target_arch = "wasm32"))] pub mod quic_flow; -pub mod retained; pub mod state; pub mod subscription; @@ -14,7 +13,5 @@ pub use limits::{ExpiringMessage, LimitsConfig, LimitsManager}; pub use queue::{MessageQueue, QueueResult, QueueStats, QueuedMessage}; #[cfg(not(target_arch = "wasm32"))] pub use quic_flow::{FlowRegistry, FlowState, FlowType}; -#[allow(deprecated)] -pub use retained::{RetainedMessage, RetainedMessageStore}; pub use state::{SessionConfig, SessionState, SessionStats}; pub use subscription::{Subscription, SubscriptionManager}; diff --git a/crates/mqtt5/src/session/flow_control.rs b/crates/mqtt5/src/session/flow_control.rs index 03299e78..f97e25b8 100644 --- a/crates/mqtt5/src/session/flow_control.rs +++ b/crates/mqtt5/src/session/flow_control.rs @@ -1,6 +1,8 @@ use crate::error::{MqttError, Result}; use crate::time::{Duration, Instant}; +use parking_lot::Mutex; use std::collections::{HashMap, VecDeque}; +use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::Arc; use tokio::sync::{Notify, RwLock, Semaphore}; @@ -25,6 +27,8 @@ pub struct FlowControlManager { inbound_receive_maximum: u16, /// Currently in-flight inbound messages from server (`packet_id` -> timestamp) inbound_in_flight: Arc>>, + quota_debt: Arc, + replay_slots: Arc>>>, } /// A pending publish request waiting for quota @@ -63,6 +67,8 @@ impl FlowControlManager { config, inbound_receive_maximum: 65535, inbound_in_flight: Arc::new(RwLock::new(HashMap::new())), + quota_debt: Arc::new(AtomicUsize::new(0)), + replay_slots: Arc::new(Mutex::new(None)), } } @@ -83,6 +89,9 @@ impl FlowControlManager { } let mut inbound = self.inbound_in_flight.write().await; + if inbound.contains_key(&packet_id) { + return Ok(()); + } if inbound.len() >= usize::from(self.inbound_receive_maximum) { return Err(MqttError::ReceiveMaximumExceeded); } @@ -100,14 +109,140 @@ impl FlowControlManager { self.inbound_in_flight.read().await.len() } + pub async fn clear_inbound(&self) { + self.inbound_in_flight.write().await.clear(); + } + + pub async fn reset_for_connection( + &mut self, + receive_maximum: u16, + retained_in_flight: &[u16], + replay: bool, + ) -> Option> { + let now = Instant::now(); + let mut in_flight = self.in_flight.write().await; + in_flight.clear(); + in_flight.extend(retained_in_flight.iter().map(|id| (*id, now))); + let held = in_flight.len(); + drop(in_flight); + + self.receive_maximum = receive_maximum; + let capacity = if receive_maximum == 0 { + Semaphore::MAX_PERMITS + } else { + usize::from(receive_maximum).saturating_sub(held) + }; + let debt = if receive_maximum == 0 { + 0 + } else { + held.saturating_sub(usize::from(receive_maximum)) + }; + self.quota_debt.store(debt, Ordering::SeqCst); + + let (main_permits, replay_slots) = if replay { + (0, Some(Arc::new(Semaphore::new(capacity)))) + } else { + (capacity, None) + }; + let previous_slots = + std::mem::replace(&mut *self.replay_slots.lock(), replay_slots.clone()); + if let Some(slots) = previous_slots { + slots.close(); + } + let previous = std::mem::replace( + &mut self.quota_semaphore, + Arc::new(Semaphore::new(main_permits)), + ); + previous.close(); + self.quota_available.notify_waiters(); + replay_slots + } + + pub async fn finish_replay(&self, slots: &Arc) { + let in_flight = self.in_flight.write().await; + let mut current = self.replay_slots.lock(); + if current + .as_ref() + .is_some_and(|active| Arc::ptr_eq(active, slots)) + { + *current = None; + let free = slots.available_permits(); + slots.close(); + self.quota_semaphore.add_permits(free); + self.quota_available.notify_waiters(); + } + drop(current); + drop(in_flight); + } + + pub async fn claim_send_quota(&self, semaphore: &Arc, packet_id: u16) -> bool { + let issued_by_current = Arc::ptr_eq(semaphore, &self.quota_semaphore) + || self + .replay_slots + .lock() + .as_ref() + .is_some_and(|slots| Arc::ptr_eq(slots, semaphore)); + if issued_by_current { + self.in_flight + .write() + .await + .insert(packet_id, Instant::now()); + } + issued_by_current + } + + /// # Errors + /// + /// Returns `FlowControlExceeded` when the backpressure timeout elapses and + /// `NotConnected` when the quota was closed by a disconnect. + pub async fn acquire_shared_send_quota(flow: &Arc>, packet_id: u16) -> Result<()> { + loop { + let (semaphore, timeout) = { + let manager = flow.read().await; + if manager.receive_maximum == 0 { + return Ok(()); + } + ( + Arc::clone(&manager.quota_semaphore), + manager.config.backpressure_timeout, + ) + }; + let acquired = match timeout { + Some(limit) => tokio::time::timeout(limit, semaphore.acquire()) + .await + .map_err(|_| MqttError::FlowControlExceeded)?, + None => semaphore.acquire().await, + }; + if let Ok(permit) = acquired { + permit.forget(); + if flow + .read() + .await + .claim_send_quota(&semaphore, packet_id) + .await + { + return Ok(()); + } + } else if Arc::ptr_eq(&semaphore, &flow.read().await.quota_semaphore) { + return Err(MqttError::NotConnected); + } + } + } + + pub fn close_send_quota(&self) { + self.quota_semaphore.close(); + if let Some(slots) = self.replay_slots.lock().take() { + slots.close(); + } + } + /// Checks if we can send a new `QoS` 1/2 message #[must_use] pub fn can_send(&self) -> bool { if self.receive_maximum == 0 { - return true; // 0 means unlimited + return true; } - // Check if we have available permits self.quota_semaphore.available_permits() > 0 } @@ -118,10 +253,9 @@ impl FlowControlManager { /// Returns an error if the operation fails pub async fn acquire_send_quota(&self, packet_id: u16) -> Result<()> { if self.receive_maximum == 0 { - return Ok(()); // Unlimited + return Ok(()); } - // Try to acquire a permit let permit_result = if let Some(timeout) = self.config.backpressure_timeout { tokio::time::timeout(timeout, self.quota_semaphore.acquire()) .await @@ -132,13 +266,11 @@ impl FlowControlManager { let permit = permit_result.map_err(|_| MqttError::FlowControlExceeded)?; - // Record the in-flight message { let mut in_flight = self.in_flight.write().await; in_flight.insert(packet_id, Instant::now()); } - // Forget the permit (keep it acquired) permit.forget(); Ok(()) @@ -151,22 +283,19 @@ impl FlowControlManager { /// Returns an error if the operation fails pub async fn try_acquire_send_quota(&self, packet_id: u16) -> Result<()> { if self.receive_maximum == 0 { - return Ok(()); // Unlimited + return Ok(()); } - // Try to acquire a permit without waiting let permit = self .quota_semaphore .try_acquire() .map_err(|_| MqttError::FlowControlExceeded)?; - // Record the in-flight message { let mut in_flight = self.in_flight.write().await; in_flight.insert(packet_id, Instant::now()); } - // Forget the permit (keep it acquired) permit.forget(); Ok(()) @@ -194,10 +323,19 @@ impl FlowControlManager { return Err(MqttError::PacketIdNotFound(packet_id)); } - // Release the quota by adding a permit back to the semaphore - self.quota_semaphore.add_permits(1); + let paid_debt = self + .quota_debt + .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |debt| { + debt.checked_sub(1) + }) + .is_ok(); + if !paid_debt { + match self.replay_slots.lock().as_ref() { + Some(slots) => slots.add_permits(1), + None => self.quota_semaphore.add_permits(1), + } + } - // Notify waiting requests self.quota_available.notify_one(); } @@ -220,9 +358,7 @@ impl FlowControlManager { let old_value = self.receive_maximum; self.receive_maximum = value; - // Adjust semaphore permits based on the change if value == 0 { - // Unlimited - give maximum permits let current_permits = self.quota_semaphore.available_permits(); let max_permits = tokio::sync::Semaphore::MAX_PERMITS; if current_permits < max_permits { @@ -230,8 +366,6 @@ impl FlowControlManager { .add_permits(max_permits - current_permits); } } else if old_value == 0 { - // Was unlimited, now limited - // Close the semaphore and create new one with proper permits let in_flight_count = self.in_flight.read().await.len(); let available_permits = if usize::from(value) > in_flight_count { usize::from(value) - in_flight_count @@ -240,7 +374,6 @@ impl FlowControlManager { }; self.quota_semaphore = Arc::new(Semaphore::new(available_permits)); } else { - // Both were limited, adjust the difference let current_permits = self.quota_semaphore.available_permits(); let in_flight_count = self.in_flight.read().await.len(); let target_permits = if usize::from(value) > in_flight_count { @@ -255,7 +388,6 @@ impl FlowControlManager { .add_permits(target_permits - current_permits); } std::cmp::Ordering::Less => { - // Need to reduce permits - acquire the difference and forget them let to_remove = current_permits - target_permits; for _ in 0..to_remove { if let Ok(permit) = self.quota_semaphore.try_acquire() { @@ -263,13 +395,10 @@ impl FlowControlManager { } } } - std::cmp::Ordering::Equal => { - // No change needed - } + std::cmp::Ordering::Equal => {} } } - // Notify waiting requests about quota changes self.quota_available.notify_waiters(); } @@ -319,7 +448,6 @@ impl FlowControlManager { } } - // Release quota for expired messages if released_count > 0 && self.receive_maximum > 0 { self.quota_semaphore.add_permits(released_count); self.quota_available.notify_waiters(); @@ -352,7 +480,6 @@ mod tests { assert!(fc.can_send()); - // Register some messages fc.register_send(1).await.unwrap(); fc.register_send(2).await.unwrap(); fc.register_send(3).await.unwrap(); @@ -360,10 +487,8 @@ mod tests { assert_eq!(fc.in_flight_count().await, 3); assert!(!fc.can_send()); - // Try to register another assert!(fc.register_send(4).await.is_err()); - // Acknowledge one fc.acknowledge(2).await.unwrap(); assert_eq!(fc.in_flight_count().await, 2); assert!(fc.can_send()); @@ -371,16 +496,14 @@ mod tests { #[tokio::test] async fn test_flow_control_unlimited() { - let fc = FlowControlManager::new(0); // 0 means unlimited + let fc = FlowControlManager::new(0); - // Should always be able to send assert!(fc.can_send()); - // Registration should be no-op fc.register_send(1).await.unwrap(); fc.register_send(2).await.unwrap(); - assert_eq!(fc.in_flight_count().await, 0); // Not tracked when unlimited + assert_eq!(fc.in_flight_count().await, 0); } #[tokio::test] @@ -390,12 +513,10 @@ mod tests { fc.register_send(1).await.unwrap(); fc.register_send(2).await.unwrap(); - // Sleep a bit tokio::time::sleep(crate::time::Duration::from_millis(10)).await; fc.register_send(3).await.unwrap(); - // Check expired with very short timeout let expired = fc.get_expired(crate::time::Duration::from_millis(5)).await; assert_eq!(expired.len(), 2); assert!(expired.contains(&1)); @@ -403,22 +524,116 @@ mod tests { assert!(!expired.contains(&3)); } + #[tokio::test] + async fn reset_for_connection_reclaims_quota_held_by_the_previous_connection() { + let flow = Arc::new(RwLock::new(FlowControlManager::new(1))); + flow.read().await.acquire_send_quota(1).await.unwrap(); + assert!(!flow.read().await.can_send()); + + let replay = flow.write().await.reset_for_connection(1, &[], false).await; + assert!(replay.is_none()); + assert_eq!(flow.read().await.in_flight_count().await, 0); + FlowControlManager::acquire_shared_send_quota(&flow, 2) + .await + .unwrap(); + assert_eq!(flow.read().await.in_flight_count().await, 1); + } + + #[tokio::test] + async fn replay_slots_take_released_quota_until_replay_finishes() { + let flow = Arc::new(RwLock::new(FlowControlManager::new(10))); + let slots = flow + .write() + .await + .reset_for_connection(1, &[], true) + .await + .unwrap(); + assert_eq!(flow.read().await.available_permits(), 0); + assert_eq!(slots.available_permits(), 1); + + slots.acquire().await.unwrap().forget(); + assert!(flow.read().await.claim_send_quota(&slots, 7).await); + flow.read().await.acknowledge(7).await.unwrap(); + assert_eq!(slots.available_permits(), 1); + assert_eq!(flow.read().await.available_permits(), 0); + + flow.read().await.finish_replay(&slots).await; + assert!(slots.is_closed()); + assert_eq!(flow.read().await.available_permits(), 1); + } + + #[tokio::test] + async fn retained_in_flight_above_receive_maximum_is_repaid_before_quota_returns() { + let flow = Arc::new(RwLock::new(FlowControlManager::new(10))); + flow.write() + .await + .reset_for_connection(1, &[1, 2], false) + .await; + assert_eq!(flow.read().await.available_permits(), 0); + + flow.read().await.acknowledge(1).await.unwrap(); + assert_eq!(flow.read().await.available_permits(), 0); + flow.read().await.acknowledge(2).await.unwrap(); + assert_eq!(flow.read().await.available_permits(), 1); + } + + #[tokio::test] + async fn stale_permit_from_a_reset_quota_is_not_claimed() { + let flow = Arc::new(RwLock::new(FlowControlManager::new(2))); + let stale = Arc::clone(&flow.read().await.quota_semaphore); + flow.write().await.reset_for_connection(2, &[], false).await; + assert!(stale.is_closed()); + assert!(!flow.read().await.claim_send_quota(&stale, 1).await); + } + + #[tokio::test] + async fn closed_send_quota_fails_waiting_publishers() { + let flow = Arc::new(RwLock::new(FlowControlManager::new(1))); + FlowControlManager::acquire_shared_send_quota(&flow, 1) + .await + .unwrap(); + let waiter = { + let flow = Arc::clone(&flow); + tokio::spawn( + async move { FlowControlManager::acquire_shared_send_quota(&flow, 2).await }, + ) + }; + tokio::time::sleep(Duration::from_millis(20)).await; + flow.read().await.close_send_quota(); + assert!(matches!( + waiter.await.unwrap(), + Err(MqttError::NotConnected) + )); + } + + #[tokio::test] + async fn duplicate_inbound_packet_id_is_not_counted_twice() { + let mut fc = FlowControlManager::new(10); + fc.set_inbound_receive_maximum(1); + fc.register_inbound_publish(1).await.unwrap(); + fc.register_inbound_publish(1).await.unwrap(); + assert_eq!(fc.inbound_in_flight_count().await, 1); + assert!(matches!( + fc.register_inbound_publish(2).await, + Err(MqttError::ReceiveMaximumExceeded) + )); + fc.clear_inbound().await; + fc.register_inbound_publish(2).await.unwrap(); + } + #[test] fn test_topic_alias_basic() { let mut ta = TopicAliasManager::new(10); - // Get or create alias let alias1 = ta.get_or_create_alias("topic/1").unwrap(); assert_eq!(alias1, 1); let alias2 = ta.get_or_create_alias("topic/2").unwrap(); assert_eq!(alias2, 2); - // Same topic should return same alias let alias1_again = ta.get_or_create_alias("topic/1").unwrap(); assert_eq!(alias1_again, 1); - // Check lookups assert_eq!(ta.get_topic(1), Some("topic/1")); assert_eq!(ta.get_alias("topic/1"), Some(1)); } @@ -427,15 +642,12 @@ mod tests { fn test_topic_alias_register() { let mut ta = TopicAliasManager::new(5); - // Register alias from peer ta.register_alias(3, "remote/topic").unwrap(); assert_eq!(ta.get_topic(3), Some("remote/topic")); - // Invalid alias assert!(ta.register_alias(0, "topic").is_err()); assert!(ta.register_alias(6, "topic").is_err()); - // Overwrite existing alias ta.register_alias(3, "new/topic").unwrap(); assert_eq!(ta.get_topic(3), Some("new/topic")); assert!(ta.get_alias("remote/topic").is_none()); @@ -451,7 +663,7 @@ mod tests { assert!(alias1.is_some()); assert!(alias2.is_some()); - assert!(alias3.is_none()); // Limit reached + assert!(alias3.is_none()); } #[test] diff --git a/crates/mqtt5/src/session/retained.rs b/crates/mqtt5/src/session/retained.rs deleted file mode 100644 index 31b1ba48..00000000 --- a/crates/mqtt5/src/session/retained.rs +++ /dev/null @@ -1,224 +0,0 @@ -#![allow(deprecated)] - -use crate::packet::publish::PublishPacket; -use crate::topic_matching::matches as topic_matches; -use crate::QoS; -use std::collections::HashMap; -use std::sync::Arc; -use tokio::sync::RwLock; - -/// Storage for retained messages -#[deprecated( - since = "0.31.5", - note = "session-level retained store is unused by the broker; use the broker's storage backend (broker::storage::RetainedMessage) instead. Scheduled for removal in 0.32.0." -)] -#[derive(Debug, Clone)] -pub struct RetainedMessageStore { - /// Map of topic names to retained messages - messages: Arc>>, -} - -/// A retained message -#[deprecated( - since = "0.31.5", - note = "session-level retained store is unused by the broker; use broker::storage::RetainedMessage instead. Scheduled for removal in 0.32.0." -)] -#[derive(Debug, Clone)] -pub struct RetainedMessage { - /// The topic name - pub topic: String, - /// The message payload (empty means clear retained message) - pub payload: Vec, - /// `QoS` level - pub qos: QoS, - /// Message properties - pub properties: crate::protocol::v5::properties::Properties, -} - -impl RetainedMessageStore { - /// Creates a new retained message store - #[must_use] - pub fn new() -> Self { - Self { - messages: Arc::new(RwLock::new(HashMap::new())), - } - } - - /// Stores or clears a retained message - pub async fn store(&self, topic: impl Into, message: Option) { - let mut messages = self.messages.write().await; - let topic = topic.into(); - - if let Some(msg) = message { - messages.insert(topic, msg); - } else { - messages.remove(&topic); - } - } - - /// Gets all retained messages matching a topic filter - pub async fn get_matching(&self, topic_filter: &str) -> Vec { - let messages = self.messages.read().await; - let mut matching = Vec::new(); - - for (topic, message) in messages.iter() { - if topic_matches(topic, topic_filter) { - matching.push(message.clone()); - } - } - - matching - } - - /// Gets a specific retained message by exact topic - pub async fn get(&self, topic: &str) -> Option { - let messages = self.messages.read().await; - messages.get(topic).cloned() - } - - /// Clears all retained messages - pub async fn clear_all(&self) { - let mut messages = self.messages.write().await; - messages.clear(); - } - - /// Gets the number of retained messages - pub async fn count(&self) -> usize { - let messages = self.messages.read().await; - messages.len() - } -} - -impl From<&PublishPacket> for RetainedMessage { - fn from(packet: &PublishPacket) -> Self { - Self { - topic: packet.topic_name.clone(), - payload: packet.payload.to_vec(), - qos: packet.qos, - properties: packet.properties.clone(), - } - } -} - -impl Default for RetainedMessageStore { - fn default() -> Self { - Self::new() - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::test_utils::TestMessageBuilder; - use crate::Properties; - - #[tokio::test] - async fn test_store_and_retrieve() { - let store = RetainedMessageStore::new(); - - // Store a retained message - let msg = RetainedMessage { - topic: "test/topic".to_string(), - payload: b"test payload".to_vec(), - qos: QoS::AtLeastOnce, - properties: Properties::default(), - }; - - store - .store("test/topic".to_string(), Some(msg.clone())) - .await; - - // Retrieve the message - let retrieved = store.get("test/topic").await; - assert!(retrieved.is_some()); - - let retrieved = retrieved.unwrap(); - assert_eq!(retrieved.topic, "test/topic"); - assert_eq!(&retrieved.payload[..], b"test payload"); - assert_eq!(retrieved.qos, QoS::AtLeastOnce); - } - - #[tokio::test] - async fn test_clear_retained_message() { - let store = RetainedMessageStore::new(); - - // Store a retained message - let msg = RetainedMessage { - topic: "test/topic".to_string(), - payload: b"test payload".to_vec(), - qos: QoS::AtMostOnce, - properties: Properties::default(), - }; - - store.store("test/topic".to_string(), Some(msg)).await; - assert_eq!(store.count().await, 1); - - // Clear the retained message - store.store("test/topic".to_string(), None).await; - assert_eq!(store.count().await, 0); - - // Verify it's gone - let retrieved = store.get("test/topic").await; - assert!(retrieved.is_none()); - } - - #[tokio::test] - async fn test_topic_matching() { - let store = RetainedMessageStore::new(); - - // Store multiple retained messages - let topics = vec![ - "home/room1/temp", - "home/room1/humidity", - "home/room2/temp", - "office/room1/temp", - ]; - - for topic in topics { - let msg = RetainedMessage { - topic: topic.to_string(), - payload: format!("data for {topic}").into_bytes(), - qos: QoS::AtMostOnce, - properties: Properties::default(), - }; - store.store(topic.to_string(), Some(msg)).await; - } - - // Test exact match - let matching = store.get_matching("home/room1/temp").await; - assert_eq!(matching.len(), 1); - assert_eq!(matching[0].topic, "home/room1/temp"); - - // Test single-level wildcard - let matching = store.get_matching("home/+/temp").await; - assert_eq!(matching.len(), 2); - - // Test multi-level wildcard - let matching = store.get_matching("home/#").await; - assert_eq!(matching.len(), 3); - - // Test no match - let matching = store.get_matching("garage/+/temp").await; - assert_eq!(matching.len(), 0); - } - - #[tokio::test] - async fn test_clear_all() { - let store = RetainedMessageStore::new(); - - // Store multiple messages - let messages = TestMessageBuilder::new() - .with_topic_prefix("topic") - .build_retained_batch(5); - - for (i, msg) in messages.into_iter().enumerate() { - store.store(format!("topic/{i}"), Some(msg)).await; - } - - assert_eq!(store.count().await, 5); - - // Clear all - store.clear_all().await; - assert_eq!(store.count().await, 0); - } -} diff --git a/crates/mqtt5/src/session/state.rs b/crates/mqtt5/src/session/state.rs index 9f58976c..da31cbd8 100644 --- a/crates/mqtt5/src/session/state.rs +++ b/crates/mqtt5/src/session/state.rs @@ -1,18 +1,18 @@ use crate::error::{MqttError, Result}; use crate::packet::publish::PublishPacket; +use crate::packet_id::PacketIdGenerator; use crate::session::flow_control::{FlowControlManager, TopicAliasManager}; use crate::session::limits::LimitsManager; use crate::session::queue::{MessageQueue, QueuedMessage}; #[cfg(not(target_arch = "wasm32"))] use crate::session::quic_flow::{FlowRegistry, FlowState}; -#[allow(deprecated)] -use crate::session::retained::{RetainedMessage, RetainedMessageStore}; use crate::session::subscription::{Subscription, SubscriptionManager}; use crate::time::{Duration, Instant}; #[cfg(not(target_arch = "wasm32"))] use crate::transport::flow::{FlowFlags, FlowId}; use crate::types::WillMessage; use std::collections::HashMap; +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; use tokio::sync::RwLock; @@ -35,7 +35,7 @@ impl Default for SessionConfig { session_expiry_interval: 0, max_queued_messages: 1000, max_queued_size: crate::constants::buffer::DEFAULT_CAPACITY - * crate::constants::buffer::DEFAULT_CAPACITY, // 1MB + * crate::constants::buffer::DEFAULT_CAPACITY, persistent: false, } } @@ -52,10 +52,9 @@ pub struct SessionState { subscriptions: Arc>, /// `QoS` 1 and 2 message queue message_queue: Arc>, - /// Unacknowledged PUBLISH packets (`packet_id` -> packet) - unacked_publishes: Arc>>, - /// Unacknowledged PUBREL packets (`packet_id` -> timestamp) - unacked_pubrels: Arc>>, + unacked_publishes: Arc>>, + unacked_pubrels: Arc>>, + outbound_send_order: AtomicU64, /// Inbound `QoS` 2 packet IDs we have PUBREC'd and owe a PUBCOMP for /// (`packet_id` -> timestamp). /// @@ -87,9 +86,6 @@ pub struct SessionState { topic_alias_out: Arc>, /// Topic alias manager for incoming messages topic_alias_in: Arc>, - /// Retained message store - #[allow(deprecated)] - retained_messages: Arc, /// Will message (to be published on abnormal disconnection) will_message: Arc>>, /// Will delay timer handle @@ -100,6 +96,12 @@ pub struct SessionState { flow_registry: Arc>, } +#[derive(Debug, Clone)] +pub enum OutboundReplay { + Publish(PublishPacket), + PubRel(u16), +} + /// The application's decision for a deferred inbound `QoS` 2 message while its handshake /// is still in flight. /// @@ -130,6 +132,7 @@ impl SessionState { config, unacked_publishes: Arc::new(RwLock::new(HashMap::new())), unacked_pubrels: Arc::new(RwLock::new(HashMap::new())), + outbound_send_order: AtomicU64::new(0), inbound_pubrecs: Arc::new(RwLock::new(HashMap::new())), inbound_delivered: Arc::new(RwLock::new(HashMap::new())), inbound_resolution: Arc::new(RwLock::new(HashMap::new())), @@ -138,13 +141,9 @@ impl SessionState { created_at: now, last_activity: Arc::new(RwLock::new(now)), clean_start, - flow_control: Arc::new(RwLock::new(FlowControlManager::new(65535))), // Default to max - topic_alias_out: Arc::new(RwLock::new(TopicAliasManager::new(0))), // Default to disabled - topic_alias_in: Arc::new(RwLock::new(TopicAliasManager::new(0))), // Default to disabled - retained_messages: Arc::new({ - #[allow(deprecated)] - RetainedMessageStore::new() - }), + flow_control: Arc::new(RwLock::new(FlowControlManager::new(65535))), + topic_alias_out: Arc::new(RwLock::new(TopicAliasManager::new(0))), + topic_alias_in: Arc::new(RwLock::new(TopicAliasManager::new(0))), will_message: Arc::new(RwLock::new(None)), will_delay_handle: Arc::new(RwLock::new(None)), limits: Arc::new(RwLock::new(LimitsManager::with_defaults())), @@ -159,6 +158,14 @@ impl SessionState { &self.client_id } + pub fn set_client_id(&mut self, client_id: String) { + self.client_id = client_id; + } + + fn next_send_order(&self) -> u64 { + self.outbound_send_order.fetch_add(1, Ordering::SeqCst) + } + #[must_use] /// Checks if this is a clean session pub fn is_clean(&self) -> bool { @@ -173,7 +180,7 @@ impl SessionState { /// Checks if session has expired pub async fn is_expired(&self) -> bool { if self.config.session_expiry_interval == 0 { - return false; // Session doesn't expire + return false; } let last_activity = *self.last_activity.read().await; @@ -274,10 +281,11 @@ impl SessionState { pub async fn store_unacked_publish(&self, packet: PublishPacket) -> Result<()> { if let Some(packet_id) = packet.packet_id { self.touch().await; + let order = self.next_send_order(); self.unacked_publishes .write() .await - .insert(packet_id, packet); + .insert(packet_id, (order, packet)); Ok(()) } else { Err(MqttError::ProtocolError( @@ -289,26 +297,29 @@ impl SessionState { /// Removes an acknowledged PUBLISH packet pub async fn remove_unacked_publish(&self, packet_id: u16) -> Option { self.touch().await; - self.unacked_publishes.write().await.remove(&packet_id) + self.unacked_publishes + .write() + .await + .remove(&packet_id) + .map(|(_, packet)| packet) } /// Gets all unacknowledged PUBLISH packets pub async fn get_unacked_publishes(&self) -> Vec { - self.unacked_publishes + let mut ordered: Vec<(u64, PublishPacket)> = self + .unacked_publishes .read() .await .values() .cloned() - .collect() + .collect(); + ordered.sort_by_key(|(order, _)| *order); + ordered.into_iter().map(|(_, packet)| packet).collect() } /// Stores an unacknowledged PUBREL packet pub async fn store_unacked_pubrel(&self, packet_id: u16) { - self.touch().await; - self.unacked_pubrels - .write() - .await - .insert(packet_id, Instant::now()); + self.store_pubrel(packet_id).await; } /// Removes an acknowledged PUBREL packet @@ -326,6 +337,49 @@ impl SessionState { self.unacked_pubrels.read().await.keys().copied().collect() } + pub async fn outbound_replay(&self) -> Vec { + let publishes = self.unacked_publishes.read().await; + let pubrels = self.unacked_pubrels.read().await; + let mut ordered: Vec<(u64, OutboundReplay)> = publishes + .values() + .map(|(order, packet)| (*order, OutboundReplay::Publish(packet.clone()))) + .chain( + pubrels + .iter() + .map(|(packet_id, (order, _))| (*order, OutboundReplay::PubRel(*packet_id))), + ) + .collect(); + drop(pubrels); + drop(publishes); + ordered.sort_by_key(|(order, _)| *order); + ordered.into_iter().map(|(_, item)| item).collect() + } + + pub async fn discard_outbound_state(&self) { + self.unacked_publishes.write().await.clear(); + self.unacked_pubrels.write().await.clear(); + } + + pub async fn complete_outbound(&self, packet_id: u16) { + self.touch().await; + self.unacked_publishes.write().await.remove(&packet_id); + self.unacked_pubrels.write().await.remove(&packet_id); + } + + pub async fn allocate_packet_id( + &self, + generator: &PacketIdGenerator, + in_use_elsewhere: impl Fn(u16) -> bool, + ) -> Option { + let publishes = self.unacked_publishes.read().await; + let pubrels = self.unacked_pubrels.read().await; + generator.next_available(|packet_id| { + publishes.contains_key(&packet_id) + || pubrels.contains_key(&packet_id) + || in_use_elsewhere(packet_id) + }) + } + /// Clears all session state pub async fn clear(&self) { self.subscriptions.write().await.clear(); @@ -480,44 +534,6 @@ impl SessionState { self.message_queue.write().await.remove_expired(timeout); } - /// Stores or clears a retained message - #[deprecated( - since = "0.31.5", - note = "session-level retained store is unused by the broker; the broker uses broker::storage::RetainedMessage. Scheduled for removal in 0.32.0." - )] - #[allow(deprecated)] - pub async fn store_retained_message(&self, packet: &PublishPacket) { - let topic = packet.topic_name.clone(); - - if packet.payload.is_empty() { - self.retained_messages.store(topic, None).await; - } else { - let message = RetainedMessage::from(packet); - self.retained_messages.store(topic, Some(message)).await; - } - } - - /// Gets retained messages matching a topic filter - #[deprecated( - since = "0.31.5", - note = "session-level retained store is unused by the broker; the broker uses broker::storage::RetainedMessage. Scheduled for removal in 0.32.0." - )] - #[allow(deprecated)] - pub async fn get_retained_messages(&self, topic_filter: &str) -> Vec { - self.retained_messages.get_matching(topic_filter).await - } - - #[must_use] - /// Gets the retained message store - #[deprecated( - since = "0.31.5", - note = "session-level retained store is unused by the broker; the broker uses broker::storage::RetainedMessage. Scheduled for removal in 0.32.0." - )] - #[allow(deprecated)] - pub fn retained_messages(&self) -> &Arc { - &self.retained_messages - } - /// Sets the Will message for this session pub async fn set_will_message(&self, will: Option) { let mut will_message = self.will_message.write().await; @@ -533,25 +549,20 @@ impl SessionState { /// Triggers Will message publication (called on abnormal disconnection) pub async fn trigger_will_message(&self) -> Option { let mut will_message = self.will_message.write().await; - let will = will_message.take(); // Remove the will message so it's only sent once + let will = will_message.take(); if let Some(ref will) = will { - // If there's a will delay interval, start the delay timer if let Some(delay_seconds) = will.properties.will_delay_interval { if delay_seconds > 0 { let delay_handle_clone = Arc::clone(&self.will_delay_handle); - // Spawn a task to handle the delay let handle = tokio::spawn(async move { tokio::time::sleep(Duration::from_secs(u64::from(delay_seconds))).await; - // The actual will message publication will be handled by the caller - // This just implements the delay }); let mut delay_handle = delay_handle_clone.write().await; *delay_handle = Some(handle); - // Return None to indicate the will should be delayed return None; } } @@ -562,11 +573,9 @@ impl SessionState { /// Cancels the Will message (called on normal disconnection) pub async fn cancel_will_message(&self) { - // Clear the will message let mut will_message = self.will_message.write().await; *will_message = None; - // Cancel any pending will delay timer let mut delay_handle = self.will_delay_handle.write().await; if let Some(handle) = delay_handle.take() { handle.abort(); @@ -579,7 +588,7 @@ impl SessionState { if let Some(ref handle) = *delay_handle { handle.is_finished() } else { - true // No delay, so it's "complete" + true } } @@ -725,17 +734,21 @@ impl SessionState { /// Store PUBREL for `QoS` 2 flow pub async fn store_pubrel(&self, packet_id: u16) { self.touch().await; + let order = self.next_send_order(); self.unacked_pubrels .write() .await - .insert(packet_id, Instant::now()); + .entry(packet_id) + .or_insert((order, Instant::now())); } - /// Complete PUBREC (after sending PUBREL) pub async fn complete_pubrec(&self, packet_id: u16) { self.touch().await; - // Remove from unacked publishes as we've moved to PUBREL phase - self.unacked_publishes.write().await.remove(&packet_id); + let mut publishes = self.unacked_publishes.write().await; + let mut pubrels = self.unacked_pubrels.write().await; + if let Some((order, _)) = publishes.remove(&packet_id) { + pubrels.insert(packet_id, (order, Instant::now())); + } } /// Complete PUBREL (after receiving PUBCOMP) @@ -862,7 +875,6 @@ pub struct SessionStats { } #[cfg(test)] -#[allow(deprecated)] mod tests { use super::*; use crate::packet::subscribe::SubscriptionOptions; @@ -882,19 +894,16 @@ mod tests { #[tokio::test] async fn test_session_expiry() { let config = SessionConfig { - session_expiry_interval: 1, // 1 second + session_expiry_interval: 1, ..Default::default() }; let session = SessionState::new("test-client".to_string(), config, false); - // Initially not expired assert!(!session.is_expired().await); - // Update last activity to past *session.last_activity.write().await = Instant::now().checked_sub(Duration::from_secs(2)).unwrap(); - // Now should be expired assert!(session.is_expired().await); } @@ -907,18 +916,15 @@ mod tests { options: SubscriptionOptions::default(), }; - // Add subscription session .add_subscription("test/topic".to_string(), sub.clone()) .await .unwrap(); - // Check matching let matches = session.matching_subscriptions("test/topic").await; assert_eq!(matches.len(), 1); assert_eq!(matches[0].0, "test/topic"); - // Remove subscription session.remove_subscription("test/topic").await.unwrap(); let matches = session.matching_subscriptions("test/topic").await; assert_eq!(matches.len(), 0); @@ -928,7 +934,6 @@ mod tests { async fn test_message_queueing() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Queue messages let msg1 = QueuedMessage { topic: "test/1".to_string(), payload: vec![1, 2, 3], @@ -950,7 +955,6 @@ mod tests { assert_eq!(session.queued_message_count().await, 2); - // Dequeue messages let messages = session.dequeue_messages(1).await; assert_eq!(messages.len(), 1); assert_eq!(session.queued_message_count().await, 1); @@ -978,7 +982,6 @@ mod tests { assert_eq!(unacked.len(), 1); assert_eq!(unacked[0].packet_id, Some(123)); - // Remove acknowledged publish let removed = session.remove_unacked_publish(123).await; assert!(removed.is_some()); assert_eq!(session.get_unacked_publishes().await.len(), 0); @@ -988,7 +991,6 @@ mod tests { async fn test_unacked_pubrel_tracking() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Store unacked pubrels session.store_unacked_pubrel(100).await; session.store_unacked_pubrel(101).await; @@ -997,7 +999,6 @@ mod tests { assert!(pubrels.contains(&100)); assert!(pubrels.contains(&101)); - // Remove acknowledged pubrel assert!(session.remove_unacked_pubrel(100).await); assert_eq!(session.get_unacked_pubrels().await.len(), 1); } @@ -1006,7 +1007,6 @@ mod tests { async fn test_session_clear() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Add some data let sub = Subscription { topic_filter: "test/#".to_string(), options: SubscriptionOptions::default(), @@ -1027,10 +1027,8 @@ mod tests { session.store_unacked_pubrel(1).await; - // Clear session session.clear().await; - // Verify everything is cleared assert_eq!(session.all_subscriptions().await.len(), 0); assert_eq!(session.queued_message_count().await, 0); assert_eq!(session.get_unacked_pubrels().await.len(), 0); @@ -1040,7 +1038,6 @@ mod tests { async fn test_session_stats() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Add some data let sub = Subscription { topic_filter: "test/#".to_string(), options: SubscriptionOptions::default(), @@ -1053,7 +1050,6 @@ mod tests { let stats = session.stats().await; assert_eq!(stats.subscription_count, 1); assert_eq!(stats.queued_message_count, 0); - // Uptime might be 0 on very fast systems, so just check it exists let _ = stats.uptime.as_nanos(); } @@ -1061,26 +1057,19 @@ mod tests { async fn test_flow_control_integration() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Set receive maximum session.set_receive_maximum(2).await; - // Should be able to send initially assert!(session.can_send_qos_message().await); - // Register in-flight messages session.register_in_flight(1).await.unwrap(); session.register_in_flight(2).await.unwrap(); - // Should not be able to send more assert!(!session.can_send_qos_message().await); - // Try to register another assert!(session.register_in_flight(3).await.is_err()); - // Acknowledge one session.acknowledge_in_flight(1).await.unwrap(); - // Should be able to send again assert!(session.can_send_qos_message().await); } @@ -1088,25 +1077,20 @@ mod tests { async fn test_topic_alias_integration() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Set topic alias maximum session.set_topic_alias_maximum_out(10).await; session.set_topic_alias_maximum_in(10).await; - // Get or create alias for outgoing let alias1 = session.get_or_create_topic_alias("topic/1").await; assert_eq!(alias1, Some(1)); - // Same topic should get same alias let alias1_again = session.get_or_create_topic_alias("topic/1").await; assert_eq!(alias1_again, Some(1)); - // Register incoming alias session .register_incoming_topic_alias(5, "incoming/topic") .await .unwrap(); - // Get topic for alias let topic = session.get_topic_for_alias(5).await; assert_eq!(topic, Some("incoming/topic".to_string())); } @@ -1114,17 +1098,15 @@ mod tests { #[tokio::test] async fn test_session_expiry_zero_interval() { let config = SessionConfig { - session_expiry_interval: 0, // Session doesn't expire + session_expiry_interval: 0, ..Default::default() }; let session = SessionState::new("test-client".to_string(), config, false); - // Update last activity to past *session.last_activity.write().await = Instant::now() .checked_sub(Duration::from_secs(100)) .unwrap(); - // Should not be expired assert!(!session.is_expired().await); } @@ -1151,7 +1133,6 @@ mod tests { .await .unwrap(); - // Check matching let matches = session.matching_subscriptions("test/foo/topic").await; assert_eq!(matches.len(), 2); @@ -1169,7 +1150,6 @@ mod tests { let session = SessionState::new("test-client".to_string(), config, true); - // Queue messages up to limit let msg1 = QueuedMessage { topic: "test/1".to_string(), payload: vec![0; 40], @@ -1197,13 +1177,10 @@ mod tests { session.queue_message(msg1).await.unwrap(); session.queue_message(msg2).await.unwrap(); - // Third message should succeed but drop the oldest message session.queue_message(msg3).await.unwrap(); - // Should still have 2 messages assert_eq!(session.queued_message_count().await, 2); - // Dequeue all and verify oldest was dropped let messages = session.dequeue_messages(3).await; assert_eq!(messages.len(), 2); assert_eq!(messages[0].topic, "test/2"); @@ -1248,13 +1225,11 @@ mod tests { session.store_unacked_publish(packet).await.unwrap(); assert_eq!(session.get_unacked_publishes().await.len(), 1); - // Complete PUBREC (removes publish, adds to pubrel) session.complete_pubrec(123).await; session.store_pubrel(123).await; assert_eq!(session.get_unacked_publishes().await.len(), 0); assert_eq!(session.get_unacked_pubrels().await.len(), 1); - // Complete PUBREL session.complete_pubrel(123).await; assert_eq!(session.get_unacked_pubrels().await.len(), 0); } @@ -1263,102 +1238,19 @@ mod tests { async fn test_packet_size_limits() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Set server maximum packet size session.set_server_maximum_packet_size(1000).await; - // Check packet size within limit assert!(session.check_packet_size(500).await.is_ok()); - // Check packet size exceeds limit assert!(session.check_packet_size(1001).await.is_err()); - // Get effective maximum assert_eq!(session.effective_maximum_packet_size().await, 1000); } - #[tokio::test] - async fn test_retained_messages() { - let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - - let packet1 = PublishPacket { - topic_name: "test/retained".to_string(), - packet_id: None, - payload: vec![1, 2, 3].into(), - qos: QoS::AtMostOnce, - retain: true, - dup: false, - properties: Properties::default(), - protocol_version: 5, - stream_id: None, - }; - - session.store_retained_message(&packet1).await; - - // Get retained messages - let retained = session.get_retained_messages("test/retained").await; - assert_eq!(retained.len(), 1); - assert_eq!(retained[0].payload, vec![1, 2, 3]); - - let packet2 = PublishPacket { - topic_name: "test/retained".to_string(), - packet_id: None, - payload: vec![].into(), - qos: QoS::AtMostOnce, - retain: true, - dup: false, - properties: Properties::default(), - protocol_version: 5, - stream_id: None, - }; - - session.store_retained_message(&packet2).await; - - // Should be cleared - let retained = session.get_retained_messages("test/retained").await; - assert_eq!(retained.len(), 0); - } - - #[tokio::test] - async fn test_retained_message_wildcard_matching() { - let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - - let packet1 = PublishPacket { - topic_name: "test/device1/status".to_string(), - packet_id: None, - payload: b"online".to_vec().into(), - qos: QoS::AtMostOnce, - retain: true, - dup: false, - properties: Properties::default(), - protocol_version: 5, - stream_id: None, - }; - - let packet2 = PublishPacket { - topic_name: "test/device2/status".to_string(), - packet_id: None, - payload: b"offline".to_vec().into(), - qos: QoS::AtMostOnce, - retain: true, - dup: false, - properties: Properties::default(), - protocol_version: 5, - stream_id: None, - }; - - session.store_retained_message(&packet1).await; - session.store_retained_message(&packet2).await; - - // Get with wildcard - let retained = session.get_retained_messages("test/+/status").await; - assert_eq!(retained.len(), 2); - } - #[tokio::test] async fn test_will_message() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Set will message let will = WillMessage { topic: "test/will".to_string(), payload: b"disconnected".to_vec(), @@ -1369,16 +1261,13 @@ mod tests { session.set_will_message(Some(will.clone())).await; - // Get will message let stored_will = session.will_message().await; assert!(stored_will.is_some()); assert_eq!(stored_will.unwrap().topic, "test/will"); - // Trigger will message (abnormal disconnection) let triggered = session.trigger_will_message().await; assert!(triggered.is_some()); - // Will should be cleared after triggering assert!(session.will_message().await.is_none()); } @@ -1386,7 +1275,6 @@ mod tests { async fn test_will_message_cancellation() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Set will message let will = WillMessage { topic: "test/will".to_string(), payload: b"disconnected".to_vec(), @@ -1397,10 +1285,8 @@ mod tests { session.set_will_message(Some(will)).await; - // Cancel will message (normal disconnection) session.cancel_will_message().await; - // Will should be cleared assert!(session.will_message().await.is_none()); } @@ -1408,9 +1294,8 @@ mod tests { async fn test_will_delay() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); - // Set will message with delay let will_props = WillProperties { - will_delay_interval: Some(1), // 1 second delay + will_delay_interval: Some(1), ..Default::default() }; @@ -1424,17 +1309,13 @@ mod tests { session.set_will_message(Some(will)).await; - // Trigger will with delay let triggered = session.trigger_will_message().await; - assert!(triggered.is_none()); // Should return None due to delay + assert!(triggered.is_none()); - // Check delay is not complete yet assert!(!session.is_will_delay_complete().await); - // Wait for delay tokio::time::sleep(Duration::from_millis(1100)).await; - // Now delay should be complete assert!(session.is_will_delay_complete().await); } @@ -1444,10 +1325,8 @@ mod tests { let initial_activity = *session.last_activity.read().await; - // Wait a bit tokio::time::sleep(Duration::from_millis(10)).await; - // Touch should update activity session.touch().await; let new_activity = *session.last_activity.read().await; @@ -1461,7 +1340,6 @@ mod tests { let initial_activity = *session.last_activity.read().await; tokio::time::sleep(Duration::from_millis(10)).await; - // Various operations should update activity let sub = Subscription { topic_filter: "test".to_string(), options: SubscriptionOptions::default(), @@ -1498,6 +1376,78 @@ mod tests { assert_eq!(session.get_unacked_publishes().await.len(), 0); } + fn outbound_publish(packet_id: u16, qos: QoS) -> PublishPacket { + PublishPacket { + topic_name: "test/topic".to_string(), + packet_id: Some(packet_id), + payload: vec![1].into(), + qos, + retain: false, + dup: false, + properties: Properties::default(), + protocol_version: 5, + stream_id: None, + } + } + + fn replay_ids(items: &[OutboundReplay]) -> Vec<(char, u16)> { + items + .iter() + .map(|item| match item { + OutboundReplay::Publish(p) => ('P', p.packet_id.unwrap()), + OutboundReplay::PubRel(id) => ('R', *id), + }) + .collect() + } + + #[tokio::test] + async fn outbound_replay_keeps_original_send_order_across_pubrec() { + let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); + for (id, qos) in [ + (30, QoS::ExactlyOnce), + (10, QoS::AtLeastOnce), + (20, QoS::ExactlyOnce), + ] { + session + .store_unacked_publish(outbound_publish(id, qos)) + .await + .unwrap(); + } + session.complete_pubrec(30).await; + session.store_pubrel(30).await; + + assert_eq!( + replay_ids(&session.outbound_replay().await), + vec![('R', 30), ('P', 10), ('P', 20)] + ); + + session.complete_outbound(10).await; + session.complete_outbound(30).await; + assert_eq!( + replay_ids(&session.outbound_replay().await), + vec![('P', 20)] + ); + + session.discard_outbound_state().await; + assert!(session.outbound_replay().await.is_empty()); + } + + #[tokio::test] + async fn allocate_packet_id_skips_ids_held_by_outbound_exchanges() { + let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); + let generator = PacketIdGenerator::new(); + session + .store_unacked_publish(outbound_publish(1, QoS::AtLeastOnce)) + .await + .unwrap(); + session.store_pubrel(2).await; + + assert_eq!( + session.allocate_packet_id(&generator, |id| id == 3).await, + Some(4) + ); + } + #[tokio::test] async fn clear_all_inbound_state_wipes_dedup_and_reports_presence() { let session = SessionState::new("test-client".to_string(), SessionConfig::default(), true); diff --git a/crates/mqtt5/src/test_utils.rs b/crates/mqtt5/src/test_utils.rs index f0bcefa2..e4b6a918 100644 --- a/crates/mqtt5/src/test_utils.rs +++ b/crates/mqtt5/src/test_utils.rs @@ -9,8 +9,6 @@ use crate::packet::Packet; use crate::protocol::v5::properties::Properties; use crate::protocol::v5::reason_codes::ReasonCode; use crate::session::limits::{ExpiringMessage, LimitsManager}; -#[allow(deprecated)] -use crate::session::retained::RetainedMessage; use crate::time::Duration; use crate::{MqttClient, QoS, Result}; use bytes::BytesMut; @@ -216,8 +214,6 @@ macro_rules! test_timeout { }; } -// ===== Client Test Utilities ===== - /// Generates a unique client ID for testing #[must_use] pub fn test_client_id(base: &str) -> String { @@ -295,18 +291,6 @@ pub fn test_expiring_message(index: u8) -> ExpiringMessage { ) } -/// Creates a test retained message with standard defaults -#[must_use] -#[allow(deprecated)] -pub fn test_retained_message(index: u8) -> RetainedMessage { - RetainedMessage { - topic: format!("topic/{index}"), - payload: vec![index], - qos: QoS::AtMostOnce, - properties: Properties::default(), - } -} - /// Builder for creating batches of test messages pub struct TestMessageBuilder { topic_prefix: String, @@ -363,20 +347,6 @@ impl TestMessageBuilder { }) .collect() } - - /// Builds a batch of retained messages - #[must_use] - #[allow(deprecated)] - pub fn build_retained_batch(self, count: u8) -> Vec { - (0..count) - .map(|i| RetainedMessage { - topic: format!("{}/{i}", self.topic_prefix), - payload: vec![i], - qos: self.qos, - properties: Properties::default(), - }) - .collect() - } } impl Default for TestMessageBuilder { @@ -464,13 +434,11 @@ mod tests { let encoded = encode_packet(&original).unwrap(); assert!(!encoded.is_empty()); - // Verify fixed header - CONNECT is packet type 1 assert_eq!(encoded[0] >> 4, 1); } #[tokio::test] async fn test_timeout_helper() { - // Should complete let result = run_with_timeout(Duration::from_millis(100), async { tokio::time::sleep(Duration::from_millis(10)).await; 42 @@ -478,7 +446,6 @@ mod tests { .await; assert_eq!(result, 42); - // Should timeout let result = std::panic::catch_unwind(|| { tokio::runtime::Runtime::new().unwrap().block_on(async { run_with_timeout(Duration::from_millis(10), async { diff --git a/crates/mqtt5/src/transport/packet_io.rs b/crates/mqtt5/src/transport/packet_io.rs index 9616aef8..9355c62b 100644 --- a/crates/mqtt5/src/transport/packet_io.rs +++ b/crates/mqtt5/src/transport/packet_io.rs @@ -9,7 +9,7 @@ use crate::transport::tls::{TlsReadHalf, TlsWriteHalf}; use crate::Transport; use bytes::{Buf, BufMut, BytesMut}; use std::future::Future; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt}; use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf}; /// Extension trait for Transport to add packet I/O methods @@ -141,22 +141,18 @@ fn encode_packet( where F: FnOnce(&mut BytesMut) -> Result<()>, { - // Encode body first to get remaining length let mut body_buf = BytesMut::new(); encode_body(&mut body_buf)?; - // Write fixed header let byte1 = (u8::from(packet_type) << 4) | (flags & crate::constants::masks::FLAGS); buf.put_u8(byte1); encode_variable_int(buf, u32::try_from(body_buf.len()).unwrap_or(u32::MAX))?; - // Write body buf.put(body_buf); Ok(()) } -// Implement PacketIo for all types that implement Transport impl PacketIo for T {} /// Packet reader trait for split read halves @@ -349,6 +345,68 @@ pub async fn read_packet_reusing_buffer( } } +/// Decodes the next complete packet in `read_buffer`, or `None` while it is incomplete. +/// +/// # Errors +/// Returns `PacketTooLarge` if the packet exceeds `max_packet_size`, or a decode +/// error if the packet is malformed. +pub fn decode_buffered_packet( + read_buffer: &mut BytesMut, + protocol_version: u8, + max_packet_size: usize, +) -> Result> { + let Some(header_len) = fixed_header_len(read_buffer)? else { + return Ok(None); + }; + let mut header_slice: &[u8] = &read_buffer[..header_len]; + let fixed_header = FixedHeader::decode(&mut header_slice)?; + let frame_len = header_len + fixed_header.remaining_length as usize; + if frame_len > max_packet_size { + return Err(MqttError::PacketTooLarge { + size: frame_len, + max: max_packet_size, + }); + } + if read_buffer.len() < frame_len { + read_buffer.reserve(frame_len - read_buffer.len()); + return Ok(None); + } + let mut frame = read_buffer.split_to(frame_len); + frame.advance(header_len); + Packet::decode_from_body_with_version( + fixed_header.packet_type, + &fixed_header, + &mut frame, + protocol_version, + ) + .map(Some) +} + +/// Reads one packet from a byte stream, buffering partial reads in `read_buffer`. +/// +/// # Errors +/// Returns an error if the stream fails or closes, or the packet is oversized or malformed. +pub async fn read_packet_from_stream( + reader: &mut R, + protocol_version: u8, + read_buffer: &mut BytesMut, + max_packet_size: usize, +) -> Result { + loop { + if let Some(packet) = + decode_buffered_packet(read_buffer, protocol_version, max_packet_size)? + { + return Ok(packet); + } + if read_buffer.capacity() - read_buffer.len() < READ_CHUNK { + read_buffer.reserve(READ_CHUNK); + } + if reader.read_buf(read_buffer).await? == 0 { + return Err(MqttError::ClientClosed); + } + } +} + fn fixed_header_len(buf: &[u8]) -> Result> { for (index, byte) in buf.iter().enumerate().skip(1) { if (byte & crate::constants::masks::CONTINUATION_BIT) == 0 { @@ -517,6 +575,27 @@ mod tests { buf.to_vec() } + #[test] + fn buffered_decode_reassembles_split_packet_and_enforces_maximum() { + let bytes = encoded_publish("split", b"payload"); + let mut buffer = BytesMut::from(&bytes[..4]); + assert!(decode_buffered_packet(&mut buffer, 5, 1024) + .unwrap() + .is_none()); + buffer.extend_from_slice(&bytes[4..]); + assert!(matches!( + decode_buffered_packet(&mut buffer, 5, 1024).unwrap(), + Some(Packet::Publish(p)) if p.topic_name == "split" + )); + assert!(buffer.is_empty()); + + let mut oversized = BytesMut::from(&bytes[..3]); + assert!(matches!( + decode_buffered_packet(&mut oversized, 5, bytes.len() - 1), + Err(MqttError::PacketTooLarge { .. }) + )); + } + #[tokio::test] async fn read_survives_cancellation_mid_packet() { let (tx, mut transport) = ChunkTransport::new(); @@ -612,7 +691,6 @@ mod tests { let mut transport = MockTransport::new(); transport.connect().await.unwrap(); - // Inject a PINGRESP packet transport .add_incoming_data(&crate::constants::packets::PINGRESP_BYTES) .await; @@ -626,7 +704,6 @@ mod tests { let mut transport = MockTransport::new(); transport.connect().await.unwrap(); - // Inject a PINGREQ packet transport .add_incoming_data(&crate::constants::packets::PINGREQ_BYTES) .await; @@ -642,7 +719,6 @@ mod tests { let mut transport = MockTransport::new(); transport.connect().await.unwrap(); - // Create a CONNACK packet using proper encoding let connack = ConnAckPacket { protocol_version: 5, session_present: false, @@ -669,25 +745,19 @@ mod tests { let mut transport = MockTransport::new(); transport.connect().await.unwrap(); - // Create a PUBLISH packet with QoS 0 let topic = "test/topic"; let payload = b"Hello MQTT"; - // Use proper encoding let mut buf = BytesMut::new(); - // Encode topic string using the proper function crate::encoding::encode_string(&mut buf, topic).unwrap(); - // Properties length (0 for no properties) buf.put_u8(0x00); - // Payload buf.extend_from_slice(payload); - // Now create the full packet with fixed header let mut data = BytesMut::new(); - data.put_u8(0x30); // PUBLISH with QoS 0 + data.put_u8(0x30); crate::encoding::encode_variable_int(&mut data, u32::try_from(buf.len()).unwrap()).unwrap(); data.extend_from_slice(&buf); @@ -710,11 +780,8 @@ mod tests { let mut transport = MockTransport::new(); transport.connect().await.unwrap(); - // Create packet with invalid remaining length (5 bytes with continuation bit) - // This must be manually constructed as it's testing invalid encoding let mut data = BytesMut::new(); data.put_u8(crate::constants::fixed_header::PUBLISH_BASE); - // Invalid variable byte integer - 5 bytes all with continuation bit data.extend_from_slice(&[0xFF, 0xFF, 0xFF, 0xFF, 0xFF]); transport.add_incoming_data(&data).await; @@ -732,7 +799,6 @@ mod tests { async fn test_read_packet_connection_closed() { let mut transport = MockTransport::new(); - // Don't add any data - read should return 0 let result = transport.read_packet(5).await; assert!(result.is_err()); } @@ -745,7 +811,7 @@ mod tests { transport.write_packet(Packet::PingReq).await.unwrap(); let written = transport.get_written_data().await; - assert_eq!(written, crate::constants::packets::PINGREQ_BYTES.to_vec()); // PINGREQ packet + assert_eq!(written, crate::constants::packets::PINGREQ_BYTES.to_vec()); } #[tokio::test] @@ -772,12 +838,10 @@ mod tests { let written = transport.get_written_data().await; - // Verify fixed header assert_eq!(written[0] >> 4, u8::from(PacketType::Publish)); - assert_eq!(written[0] & crate::constants::masks::FLAGS, 0x02); // QoS 1 flag + assert_eq!(written[0] & crate::constants::masks::FLAGS, 0x02); - // Should contain topic, packet ID, and payload - assert!(written.len() > 2 + 4 + 2 + 3); // header + topic + packet_id + payload + assert!(written.len() > 2 + 4 + 2 + 3); } #[tokio::test] @@ -807,14 +871,12 @@ mod tests { let written = transport.get_written_data().await; - // Verify fixed header - assert_eq!(written[0], 0x82); // SUBSCRIBE with required flags - assert!(written.len() > 2); // Has content + assert_eq!(written[0], 0x82); + assert!(written.len() > 2); } #[tokio::test] async fn test_roundtrip_packets() { - // Test that we can write and read back various packet types let test_packets = vec![ Packet::PingReq, Packet::PingResp, @@ -839,7 +901,6 @@ mod tests { let read_packet = read_transport.read_packet(5).await.unwrap(); - // Basic type check match (&packet, &read_packet) { (Packet::PingReq, Packet::PingReq) | (Packet::PingResp, Packet::PingResp) => {} (Packet::ConnAck(a), Packet::ConnAck(b)) => { @@ -855,12 +916,11 @@ mod tests { async fn test_encode_packet_helper() { let mut buf = BytesMut::new(); - // Test encoding a simple packet encode_packet(&mut buf, PacketType::PingReq, 0, |_| Ok(())).unwrap(); assert_eq!(buf.len(), 2); - assert_eq!(buf[0], crate::constants::fixed_header::PINGREQ); // PINGREQ type - assert_eq!(buf[1], 0x00); // Zero length + assert_eq!(buf[0], crate::constants::fixed_header::PINGREQ); + assert_eq!(buf[1], 0x00); } #[tokio::test] @@ -868,7 +928,6 @@ mod tests { let mut transport = MockTransport::new(); transport.connect().await.unwrap(); - // Create a publish with large payload to test variable length encoding let mut large_payload = vec![0u8; 200]; for (i, byte) in large_payload.iter_mut().enumerate() { *byte = u8::try_from(i % 256).expect("modulo 256 always fits in u8"); @@ -893,8 +952,7 @@ mod tests { let written = transport.get_written_data().await; - // Verify the remaining length uses 2 bytes (since payload > 127) - assert!(written[1] & crate::constants::masks::CONTINUATION_BIT != 0); // Continuation bit set - assert!(written.len() > 200); // Contains the large payload + assert!(written[1] & crate::constants::masks::CONTINUATION_BIT != 0); + assert!(written.len() > 200); } } diff --git a/crates/mqtt5/src/transport/websocket.rs b/crates/mqtt5/src/transport/websocket.rs index 8ac66e23..d73ec72b 100644 --- a/crates/mqtt5/src/transport/websocket.rs +++ b/crates/mqtt5/src/transport/websocket.rs @@ -42,9 +42,10 @@ use crate::error::{MqttError, Result}; use crate::packet::Packet; use crate::time::Duration; -use crate::transport::packet_io::{PacketReader, PacketWriter}; +use crate::transport::packet_io::{decode_buffered_packet, PacketReader, PacketWriter}; use crate::transport::tls::TlsConfig; use crate::Transport; +use bytes::{Buf, Bytes, BytesMut}; use futures_util::{stream::SplitSink, stream::SplitStream, StreamExt}; use std::collections::HashMap; use std::net::SocketAddr; @@ -71,9 +72,6 @@ pub struct WebSocketConfig { pub user_agent: Option, /// TLS configuration for secure WebSocket connections (wss://) pub tls_config: Option, - /// Whether to verify TLS certificates (for wss://) - deprecated, use `tls_config` - #[deprecated(note = "Use tls_config field instead")] - pub verify_tls: bool, } impl WebSocketConfig { @@ -102,8 +100,6 @@ impl WebSocketConfig { headers: HashMap::new(), user_agent: Some("mqtt-v5/0.4.0".to_string()), tls_config: None, - #[allow(deprecated)] - verify_tls: true, }) } @@ -145,21 +141,6 @@ impl WebSocketConfig { self } - /// Sets whether to verify TLS certificates for wss:// connections - /// - /// # Safety - /// - /// Disabling TLS verification is insecure and should only be used for testing - #[deprecated(note = "Use with_tls_config instead")] - #[must_use] - pub fn with_tls_verification(mut self, verify: bool) -> Self { - #[allow(deprecated)] - { - self.verify_tls = verify; - } - self - } - /// Sets a custom TLS configuration for wss:// connections #[must_use] pub fn with_tls_config(mut self, tls_config: TlsConfig) -> Self { @@ -216,12 +197,10 @@ impl WebSocketConfig { )); } - // Create TLS config if it doesn't exist if self.tls_config.is_none() { self = self.with_tls_auto()?; } - // Add client certificate to TLS config if let Some(ref mut tls_config) = self.tls_config { tls_config.load_client_cert_pem(cert_path)?; tls_config.load_client_key_pem(key_path)?; @@ -245,12 +224,10 @@ impl WebSocketConfig { )); } - // Create TLS config if it doesn't exist if self.tls_config.is_none() { self = self.with_tls_auto()?; } - // Add client certificate to TLS config if let Some(ref mut tls_config) = self.tls_config { tls_config.load_client_cert_pem_bytes(cert_pem)?; tls_config.load_client_key_pem_bytes(key_pem)?; @@ -274,12 +251,10 @@ impl WebSocketConfig { )); } - // Create TLS config if it doesn't exist if self.tls_config.is_none() { self = self.with_tls_auto()?; } - // Add CA certificate to TLS config if let Some(ref mut tls_config) = self.tls_config { tls_config.load_ca_cert_pem(ca_path)?; } @@ -302,12 +277,10 @@ impl WebSocketConfig { )); } - // Create TLS config if it doesn't exist if self.tls_config.is_none() { self = self.with_tls_auto()?; } - // Add CA certificate to TLS config if let Some(ref mut tls_config) = self.tls_config { tls_config.load_ca_cert_pem_bytes(ca_pem)?; } @@ -384,7 +357,6 @@ impl WebSocketTransport { /// Gets the negotiated subprotocol (if any) #[must_use] pub fn subprotocol(&self) -> Option<&str> { - // In a real implementation, this would return the negotiated subprotocol self.config.subprotocols.first().map(String::as_str) } @@ -401,7 +373,10 @@ impl WebSocketTransport { let connection = self.connection.ok_or(MqttError::NotConnected)?; let (write, read) = connection.split(); - let read_handle = WebSocketReadHandle { reader: read }; + let read_handle = WebSocketReadHandle { + reader: read, + buffer: BytesMut::from(&self.read_buffer[..]), + }; let write_handle = WebSocketWriteHandle { writer: write }; Ok((read_handle, write_handle)) @@ -411,6 +386,7 @@ impl WebSocketTransport { /// WebSocket read handle for split operations pub struct WebSocketReadHandle { reader: SplitStream>>, + buffer: BytesMut, } /// WebSocket write handle for split operations @@ -419,29 +395,73 @@ pub struct WebSocketWriteHandle { } impl WebSocketReadHandle { - /// Reads data from the WebSocket. - /// - /// # Errors - /// Returns an error if the connection is closed or a read error occurs. - pub async fn read(&mut self, buf: &mut [u8]) -> Result { + async fn next_binary(&mut self) -> Result { loop { match self.reader.next().await { - Some(Ok(Message::Binary(data))) => { - let len = data.len().min(buf.len()); - buf[..len].copy_from_slice(&data[..len]); - return Ok(len); + Some(Ok(Message::Binary(data))) => return Ok(data), + Some(Ok(Message::Text(_))) => { + return Err(MqttError::ProtocolError( + "WebSocket text frame received [MQTT-6.0.0-1]".to_string(), + )) } Some(Ok(Message::Close(_))) | None => return Err(MqttError::ClientClosed), - Some(Ok( - Message::Ping(_) | Message::Pong(_) | Message::Text(_) | Message::Frame(_), - )) => {} + Some(Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_))) => {} Some(Err(e)) => return Err(MqttError::Io(e.to_string())), } } } + + /// Reads data from the WebSocket. + /// + /// # Errors + /// Returns an error if the connection is closed, a text frame arrives, or a read error occurs. + pub async fn read(&mut self, buf: &mut [u8]) -> Result { + if self.buffer.is_empty() { + let data = self.next_binary().await?; + self.buffer.extend_from_slice(&data); + } + let len = self.buffer.len().min(buf.len()); + buf[..len].copy_from_slice(&self.buffer[..len]); + self.buffer.advance(len); + Ok(len) + } + + /// Reads one MQTT packet, reassembling it from the binary frame byte stream + /// regardless of frame boundaries [MQTT-6.0.0-2]. + /// + /// # Errors + /// Returns an error if the connection fails or closes, a text frame arrives, + /// or the packet is oversized or malformed. + pub async fn read_packet_limited( + &mut self, + protocol_version: u8, + max_packet_size: usize, + ) -> Result { + loop { + if let Some(packet) = + decode_buffered_packet(&mut self.buffer, protocol_version, max_packet_size)? + { + return Ok(packet); + } + let data = self.next_binary().await?; + self.buffer.extend_from_slice(&data); + } + } } impl WebSocketWriteHandle { + /// Sends a WebSocket Close frame and closes the sink. + /// + /// # Errors + /// Returns an error if the close handshake cannot be written. + pub async fn close(&mut self) -> Result<()> { + use futures_util::SinkExt; + self.writer + .close() + .await + .map_err(|e| MqttError::Io(e.to_string())) + } + /// Writes data to the WebSocket. /// /// # Errors @@ -457,29 +477,7 @@ impl WebSocketWriteHandle { impl PacketReader for WebSocketReadHandle { async fn read_packet(&mut self, protocol_version: u8) -> Result { - use crate::packet::FixedHeader; - use bytes::BytesMut; - use futures_util::StreamExt; - - match self.reader.next().await { - Some(Ok(Message::Binary(data))) => { - let mut buf = BytesMut::from(&data[..]); - - let fixed_header = FixedHeader::decode(&mut buf)?; - - Packet::decode_from_body_with_version( - fixed_header.packet_type, - &fixed_header, - &mut buf, - protocol_version, - ) - } - Some(Ok(Message::Close(_))) | None => Err(MqttError::ClientClosed), - Some(Ok(_)) => Err(MqttError::ProtocolError( - "Unexpected WebSocket message type".to_string(), - )), - Some(Err(e)) => Err(MqttError::Io(e.to_string())), - } + self.read_packet_limited(protocol_version, usize::MAX).await } } @@ -491,7 +489,6 @@ impl PacketWriter for WebSocketWriteHandle { let mut buf = BytesMut::with_capacity(1024); crate::transport::packet_io::encode_packet_to_buffer(&packet, &mut buf)?; - // Send as WebSocket binary frame self.writer .send(Message::Binary(buf.to_vec().into())) .await @@ -613,9 +610,13 @@ impl Transport for WebSocketTransport { debug!("WebSocket connection closed by remote"); return Err(MqttError::ClientClosed); } - Some(Ok( - Message::Ping(_) | Message::Pong(_) | Message::Text(_) | Message::Frame(_), - )) => {} + Some(Ok(Message::Text(_))) => { + self.connected = false; + return Err(MqttError::ProtocolError( + "WebSocket text frame received [MQTT-6.0.0-1]".to_string(), + )); + } + Some(Ok(Message::Ping(_) | Message::Pong(_) | Message::Frame(_))) => {} Some(Err(e)) => { self.connected = false; return Err(MqttError::Io(e.to_string())); @@ -732,7 +733,7 @@ mod tests { assert_eq!(config.url.as_str(), "wss://broker.example.com/mqtt"); assert!(config.is_secure()); assert_eq!(config.host(), Some("broker.example.com")); - assert_eq!(config.port(), 443); // Default HTTPS port + assert_eq!(config.port(), 443); } #[test] @@ -780,13 +781,10 @@ mod tests { assert!(!transport.is_connected()); - // Connection will fail since there's no WebSocket server at localhost:59999, - // but this tests that the connect method works as expected let result = transport.connect().await; assert!(result.is_err()); assert!(!transport.is_connected()); - // Should fail to connect again (already failed state) let result = transport.connect().await; assert!(result.is_err()); } @@ -800,7 +798,6 @@ mod tests { assert!(transport.read(&mut buf).await.is_err()); assert!(transport.write(b"test").await.is_err()); - // Close should succeed even when not connected assert!(transport.close().await.is_ok()); } @@ -809,10 +806,8 @@ mod tests { let config = WebSocketConfig::new("ws://localhost:8080/mqtt").unwrap(); let mut transport = WebSocketTransport::new(config); - // Connection will fail, but we can still test the close method let _result = transport.connect().await; - // Close should work even if not connected transport.close().await.unwrap(); assert!(!transport.is_connected()); } @@ -831,7 +826,6 @@ mod tests { #[test] fn test_websocket_config_tls_auto() { - // Should work for wss:// with IP address let config = WebSocketConfig::new("wss://127.0.0.1:8443/mqtt") .unwrap() .with_tls_auto() @@ -842,13 +836,11 @@ mod tests { assert_eq!(tls_config.addr.port(), 8443); assert_eq!(tls_config.hostname, "127.0.0.1"); - // Should fail for ws:// let result = WebSocketConfig::new("ws://127.0.0.1:8080/mqtt") .unwrap() .with_tls_auto(); assert!(result.is_err()); - // Test with default port let config_default = WebSocketConfig::new("wss://127.0.0.1/mqtt") .unwrap() .with_tls_auto() @@ -873,7 +865,6 @@ mod tests { assert!(tls_config.client_cert.is_some()); assert!(tls_config.client_key.is_some()); - // Should fail for ws:// let result = WebSocketConfig::new("ws://127.0.0.1/mqtt") .unwrap() .with_client_auth_from_bytes(cert_pem, key_pem); @@ -893,7 +884,6 @@ mod tests { let tls_config = config.tls_config().unwrap(); assert!(tls_config.root_certs.is_some()); - // Should fail for ws:// let result = WebSocketConfig::new("ws://127.0.0.1/mqtt") .unwrap() .with_ca_cert_from_bytes(ca_pem); @@ -928,7 +918,7 @@ mod tests { let tls_config = config.take_tls_config(); assert!(tls_config.is_some()); - assert!(config.tls_config().is_none()); // Should be None after taking + assert!(config.tls_config().is_none()); let tls_config = tls_config.unwrap(); assert_eq!(tls_config.hostname, "127.0.0.1"); diff --git a/crates/mqtt5/src/types.rs b/crates/mqtt5/src/types.rs index 374315d5..67210f02 100644 --- a/crates/mqtt5/src/types.rs +++ b/crates/mqtt5/src/types.rs @@ -19,6 +19,22 @@ pub struct ConnectOptions { /// an `AckToken` the application resolves after durable processing. Requires a /// persistent session and a bounded receive maximum; see `validate_deferred_ack`. pub deferred_ack: bool, + /// Accept a broker-held session that this client instance has no local state for. + /// + /// By default a client that holds no session state (a fresh `MqttClient` that has not + /// yet connected, or any `Clean Start = 1` connection) treats a CONNACK with + /// `Session Present = 1` as a protocol violation per `[MQTT-3.2.2-4]`: it sends + /// DISCONNECT with reason code 0x82 (Protocol Error), closes the network connection and + /// `connect` returns an error. + /// + /// Setting this to `true` lets a fresh client connecting with `Clean Start = 0` resume + /// the session the broker kept for its client identifier, for example after a process + /// restart. The broker's subscriptions and queued messages are resumed, but any + /// outbound or inbound `QoS` 1/2 exchanges the previous process had in flight are lost + /// locally, so delivery across the restart is at-least-once: this is the deferred-ack + /// crash-recovery pattern, where messages whose `AckToken` was never resolved are + /// redelivered to the new process. It has no effect on `Clean Start = 1` connections. + pub resume_existing_session: bool, } impl ConnectOptions { @@ -31,9 +47,19 @@ impl ConnectOptions { keepalive_config: None, codec_registry: None, deferred_ack: false, + resume_existing_session: false, } } + /// Opts in to resuming a broker-held session without local session state. + /// + /// See [`ConnectOptions::resume_existing_session`]. + #[must_use] + pub fn with_resume_existing_session(mut self, resume: bool) -> Self { + self.resume_existing_session = resume; + self + } + /// Enables deferred acknowledgement. Connection-wide: every `subscribe_with_ack` /// subscription delivers an `AckToken`. Must be paired with a persistent session /// and an explicit bounded receive maximum (see `validate_deferred_ack`). @@ -208,6 +234,7 @@ impl std::fmt::Debug for ConnectOptions { &self.codec_registry.as_ref().map(|_| "CodecRegistry"), ) .field("deferred_ack", &self.deferred_ack) + .field("resume_existing_session", &self.resume_existing_session) .finish() } } diff --git a/crates/mqtt5/tests/conf_client_a.rs b/crates/mqtt5/tests/conf_client_a.rs new file mode 100644 index 00000000..2a1919e5 --- /dev/null +++ b/crates/mqtt5/tests/conf_client_a.rs @@ -0,0 +1,1670 @@ +use mqtt5::{ + AuthHandler, AuthResponse, ConnectOptions, ConnectResult, MqttClient, MqttError, + PublishOptions, PublishProperties, QoS, +}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; +use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; +use tokio::time::timeout; + +const T: Duration = Duration::from_secs(5); + +const CONNECT: u8 = 1; +const PUBLISH: u8 = 3; +const PUBACK: u8 = 4; +const PUBREC: u8 = 5; +const PUBREL: u8 = 6; +const PUBCOMP: u8 = 7; +const SUBSCRIBE: u8 = 8; +const UNSUBSCRIBE: u8 = 10; +const PINGREQ: u8 = 12; +const DISCONNECT: u8 = 14; +const AUTH: u8 = 15; + +const P_ASSIGNED_CLIENT_ID: u8 = 0x12; +const P_SERVER_KEEP_ALIVE: u8 = 0x13; +const P_AUTH_METHOD: u8 = 0x15; +const P_AUTH_DATA: u8 = 0x16; +const P_REQUEST_PROBLEM_INFO: u8 = 0x17; +const P_REASON_STRING: u8 = 0x1F; +const P_RECEIVE_MAXIMUM: u8 = 0x21; +const P_TOPIC_ALIAS_MAXIMUM: u8 = 0x22; +const P_TOPIC_ALIAS: u8 = 0x23; +const P_MAXIMUM_QOS: u8 = 0x24; +const P_RETAIN_AVAILABLE: u8 = 0x25; +const P_USER_PROPERTY: u8 = 0x26; +const P_MAXIMUM_PACKET_SIZE: u8 = 0x27; +const P_WILDCARD_AVAILABLE: u8 = 0x28; + +#[derive(Debug, Clone)] +struct Raw { + first: u8, + body: Vec, + wire_len: usize, + remaining_len_bytes: Vec, +} + +impl Raw { + fn ptype(&self) -> u8 { + self.first >> 4 + } + fn flags(&self) -> u8 { + self.first & 0x0F + } +} + +fn varint(mut n: usize) -> Vec { + let mut out = Vec::new(); + loop { + let mut b = u8::try_from(n % 128).unwrap(); + n /= 128; + if n > 0 { + b |= 0x80; + } + out.push(b); + if n == 0 { + return out; + } + } +} + +fn read_varint(buf: &[u8], pos: &mut usize) -> usize { + let mut mult = 1usize; + let mut val = 0usize; + loop { + let b = buf[*pos]; + *pos += 1; + val += usize::from(b & 0x7F) * mult; + if b & 0x80 == 0 { + return val; + } + mult *= 128; + } +} + +fn read_u16(buf: &[u8], pos: &mut usize) -> u16 { + let v = u16::from_be_bytes([buf[*pos], buf[*pos + 1]]); + *pos += 2; + v +} + +fn read_bin(buf: &[u8], pos: &mut usize) -> Vec { + let len = usize::from(read_u16(buf, pos)); + let v = buf[*pos..*pos + len].to_vec(); + *pos += len; + v +} + +fn packet(first: u8, body: &[u8]) -> Vec { + let mut out = vec![first]; + out.extend(varint(body.len())); + out.extend_from_slice(body); + out +} + +fn p_u8(id: u8, v: u8) -> Vec { + vec![id, v] +} +fn p_u16(id: u8, v: u16) -> Vec { + let mut o = vec![id]; + o.extend(v.to_be_bytes()); + o +} +fn p_u32(id: u8, v: u32) -> Vec { + let mut o = vec![id]; + o.extend(v.to_be_bytes()); + o +} +fn mqtt_str(s: &[u8]) -> Vec { + let mut o = u16::try_from(s.len()).unwrap().to_be_bytes().to_vec(); + o.extend_from_slice(s); + o +} +fn p_str(id: u8, s: &[u8]) -> Vec { + let mut o = vec![id]; + o.extend(mqtt_str(s)); + o +} +fn p_pair(k: &[u8], v: &[u8]) -> Vec { + let mut o = vec![P_USER_PROPERTY]; + o.extend(mqtt_str(k)); + o.extend(mqtt_str(v)); + o +} + +fn with_props(props: &[u8]) -> Vec { + let mut o = varint(props.len()); + o.extend_from_slice(props); + o +} + +fn connack(session_present: bool, props: &[u8]) -> Vec { + let mut body = vec![u8::from(session_present), 0x00]; + body.extend(with_props(props)); + packet(0x20, &body) +} + +fn publish_bytes( + first: u8, + topic: &[u8], + packet_id: Option, + props: &[u8], + payload: &[u8], +) -> Vec { + let mut body = mqtt_str(topic); + if let Some(id) = packet_id { + body.extend(id.to_be_bytes()); + } + body.extend(with_props(props)); + body.extend_from_slice(payload); + packet(first, &body) +} + +fn ack_with_props(first: u8, packet_id: u16, rc: u8, props: &[u8]) -> Vec { + let mut body = packet_id.to_be_bytes().to_vec(); + body.push(rc); + body.extend(with_props(props)); + packet(first, &body) +} + +fn prop_value_len(id: u8, buf: &[u8], pos: usize) -> usize { + match id { + 0x01 | 0x17 | 0x19 | 0x24 | 0x25 | 0x28 | 0x29 | 0x2A => 1, + 0x13 | 0x21 | 0x22 | 0x23 => 2, + 0x02 | 0x11 | 0x18 | 0x27 => 4, + 0x0B => { + let mut p = pos; + let start = p; + read_varint(buf, &mut p); + p - start + } + 0x26 => { + let l1 = usize::from(u16::from_be_bytes([buf[pos], buf[pos + 1]])); + let l2 = usize::from(u16::from_be_bytes([buf[pos + 2 + l1], buf[pos + 3 + l1]])); + 4 + l1 + l2 + } + _ => 2 + usize::from(u16::from_be_bytes([buf[pos], buf[pos + 1]])), + } +} + +fn parse_props(buf: &[u8], pos: &mut usize) -> Vec<(u8, Vec)> { + let len = read_varint(buf, pos); + let end = *pos + len; + let mut out = Vec::new(); + while *pos < end { + let id = u8::try_from(read_varint(buf, pos)).unwrap(); + let vlen = prop_value_len(id, buf, *pos); + out.push((id, buf[*pos..*pos + vlen].to_vec())); + *pos += vlen; + } + out +} + +fn has_prop(props: &[(u8, Vec)], id: u8) -> bool { + props.iter().any(|(i, _)| *i == id) +} + +#[derive(Debug)] +struct ParsedConnect { + flags: u8, + keep_alive: u16, + props: Vec<(u8, Vec)>, + client_id: Vec, + will_props: Option)>>, + will_topic: Option>, + will_payload: Option>, + username: Option>, + password: Option>, + trailing: usize, +} + +fn parse_connect(raw: &Raw) -> ParsedConnect { + let b = &raw.body; + let mut pos = 0; + let name = read_bin(b, &mut pos); + assert_eq!(name, b"MQTT"); + assert_eq!(b[pos], 5); + pos += 1; + let flags = b[pos]; + pos += 1; + let keep_alive = read_u16(b, &mut pos); + let props = parse_props(b, &mut pos); + let client_id = read_bin(b, &mut pos); + let (will_props, will_topic, will_payload) = if flags & 0x04 != 0 { + let wp = parse_props(b, &mut pos); + let wt = read_bin(b, &mut pos); + let wpl = read_bin(b, &mut pos); + (Some(wp), Some(wt), Some(wpl)) + } else { + (None, None, None) + }; + let username = (flags & 0x80 != 0).then(|| read_bin(b, &mut pos)); + let password = (flags & 0x40 != 0).then(|| read_bin(b, &mut pos)); + ParsedConnect { + flags, + keep_alive, + props, + client_id, + will_props, + will_topic, + will_payload, + username, + password, + trailing: b.len() - pos, + } +} + +#[derive(Debug)] +struct ParsedPublish { + qos: u8, + retain: bool, + dup: bool, + topic: Vec, + packet_id: Option, + props: Vec<(u8, Vec)>, + payload: Vec, +} + +fn parse_publish(raw: &Raw) -> ParsedPublish { + let b = &raw.body; + let mut pos = 0; + let qos = (raw.flags() >> 1) & 0x03; + let topic = read_bin(b, &mut pos); + let packet_id = (qos > 0).then(|| read_u16(b, &mut pos)); + let props = parse_props(b, &mut pos); + ParsedPublish { + qos, + retain: raw.flags() & 0x01 != 0, + dup: raw.flags() & 0x08 != 0, + topic, + packet_id, + props, + payload: b[pos..].to_vec(), + } +} + +async fn read_raw(s: &mut TcpStream) -> Option { + let mut first = [0u8; 1]; + s.read_exact(&mut first).await.ok()?; + let mut remaining_len_bytes = Vec::new(); + let mut mult = 1usize; + let mut len = 0usize; + loop { + let mut b = [0u8; 1]; + s.read_exact(&mut b).await.ok()?; + remaining_len_bytes.push(b[0]); + len += usize::from(b[0] & 0x7F) * mult; + if b[0] & 0x80 == 0 { + break; + } + mult *= 128; + } + let mut body = vec![0u8; len]; + s.read_exact(&mut body).await.ok()?; + Some(Raw { + first: first[0], + wire_len: 1 + remaining_len_bytes.len() + len, + remaining_len_bytes, + body, + }) +} + +enum Next { + Packet(Raw), + Closed, + Silent, +} + +async fn next_packet(s: &mut TcpStream, d: Duration) -> Next { + match timeout(d, read_raw(s)).await { + Ok(Some(r)) => Next::Packet(r), + Ok(None) => Next::Closed, + Err(_) => Next::Silent, + } +} + +async fn next_non_ping(s: &mut TcpStream, d: Duration) -> Next { + let deadline = tokio::time::Instant::now() + d; + loop { + let left = deadline.saturating_duration_since(tokio::time::Instant::now()); + match next_packet(s, left).await { + Next::Packet(r) if r.ptype() == PINGREQ => { + let _ = s.write_all(&[0xD0, 0x00]).await; + } + other => return other, + } + } +} + +async fn next_of_type(s: &mut TcpStream, ptype: u8, d: Duration) -> Option { + let deadline = tokio::time::Instant::now() + d; + loop { + let left = deadline.saturating_duration_since(tokio::time::Instant::now()); + match next_packet(s, left).await { + Next::Packet(r) if r.ptype() == ptype => return Some(r), + Next::Packet(r) if r.ptype() == PINGREQ => { + let _ = s.write_all(&[0xD0, 0x00]).await; + } + Next::Packet(_) => {} + Next::Closed | Next::Silent => return None, + } + } +} + +struct CloseObservation { + closed: bool, + disconnect_rc: Option, +} + +async fn observe_close(s: &mut TcpStream, d: Duration) -> CloseObservation { + let deadline = tokio::time::Instant::now() + d; + let mut disconnect_rc = None; + loop { + let left = deadline.saturating_duration_since(tokio::time::Instant::now()); + match next_non_ping(s, left).await { + Next::Packet(r) if r.ptype() == DISCONNECT => { + disconnect_rc = Some(r.body.first().copied().unwrap_or(0)); + } + Next::Packet(_) => {} + Next::Closed => { + return CloseObservation { + closed: true, + disconnect_rc, + } + } + Next::Silent => { + return CloseObservation { + closed: false, + disconnect_rc, + } + } + } + } +} + +fn opts(id: &str) -> ConnectOptions { + ConnectOptions::new(id) + .with_automatic_reconnect(false) + .with_keep_alive(Duration::from_secs(60)) +} + +fn reconnecting_opts(id: &str) -> ConnectOptions { + ConnectOptions::new(id) + .with_automatic_reconnect(true) + .with_reconnect_delay(Duration::from_millis(100), Duration::from_millis(300)) + .with_keep_alive(Duration::from_secs(60)) +} + +struct Setup { + client: MqttClient, + listener: TcpListener, + stream: TcpStream, + connect: Raw, + result: mqtt5::Result, +} + +async fn start(options: ConnectOptions, session_present: bool, props: &[u8]) -> Setup { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("mqtt://{}", listener.local_addr().unwrap()); + let client = MqttClient::with_options(options.clone()); + let c = client.clone(); + let handle = tokio::spawn(async move { Box::pin(c.connect_with_options(&url, options)).await }); + let (mut stream, _) = timeout(T, listener.accept()).await.unwrap().unwrap(); + let connect = timeout(T, read_raw(&mut stream)).await.unwrap().unwrap(); + assert_eq!(connect.ptype(), CONNECT); + stream + .write_all(&connack(session_present, props)) + .await + .unwrap(); + let result = timeout(T, handle).await.unwrap().unwrap(); + Setup { + client, + listener, + stream, + connect, + result, + } +} + +async fn accept_next( + listener: &TcpListener, + session_present: bool, + props: &[u8], +) -> (TcpStream, Raw) { + let (mut stream, _) = timeout(Duration::from_secs(8), listener.accept()) + .await + .expect("client did not reconnect") + .unwrap(); + let connect = timeout(T, read_raw(&mut stream)).await.unwrap().unwrap(); + assert_eq!(connect.ptype(), CONNECT); + stream + .write_all(&connack(session_present, props)) + .await + .unwrap(); + (stream, connect) +} + +async fn wait_connected(client: &MqttClient, want: bool) { + for _ in 0..100 { + if client.is_connected().await == want { + return; + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + panic!("client connected state never became {want}"); +} + +async fn reconnect_with( + conn1_props: &[u8], + conn2_props: &[u8], + id: &str, +) -> (MqttClient, TcpListener, TcpStream) { + let setup = start(reconnecting_opts(id), false, conn1_props).await; + setup.result.expect("first connect"); + let Setup { + client, + listener, + stream, + .. + } = setup; + drop(stream); + wait_connected(&client, false).await; + let (stream2, _) = accept_next(&listener, false, conn2_props).await; + wait_connected(&client, true).await; + (client, listener, stream2) +} + +fn qos(q: QoS) -> PublishOptions { + PublishOptions { + qos: q, + ..Default::default() + } +} + +fn with_alias(alias: u16) -> PublishOptions { + PublishOptions { + properties: PublishProperties { + topic_alias: Some(alias), + ..Default::default() + }, + ..Default::default() + } +} + +async fn assert_no_alias_sent(client: &MqttClient, stream: &mut TcpStream, alias: u16, id: &str) { + let result = client + .publish_with_options("alias/topic", b"x".to_vec(), with_alias(alias)) + .await; + if let Some(raw) = next_of_type(stream, PUBLISH, Duration::from_secs(1)).await { + let p = parse_publish(&raw); + assert!( + !has_prop(&p.props, P_TOPIC_ALIAS), + "{id} VIOLATION: client sent Topic Alias {alias} in PUBLISH (props {:?}); publish() returned {result:?}", + p.props + ); + } +} + +#[tokio::test] +async fn mqtt_3_2_2_4_fresh_client_clean_start_1_session_present_1_must_close() { + let mut s = start(opts("sp-clean"), true, &[]).await; + let obs = observe_close(&mut s.stream, Duration::from_secs(2)).await; + assert!( + s.result.is_err() && obs.closed, + "MQTT-3.2.2-4 VIOLATION: fresh client (Clean Start=1, no session state) accepted CONNACK Session Present=1: connect result {:?}, network closed={}", + s.result, + obs.closed + ); +} + +#[tokio::test] +async fn mqtt_3_2_2_4_fresh_client_clean_start_0_session_present_1_must_close() { + let o = opts("sp-fresh") + .with_clean_start(false) + .with_session_expiry_interval(300); + let mut s = start(o, true, &[]).await; + let obs = observe_close(&mut s.stream, Duration::from_secs(2)).await; + assert!( + s.result.is_err() && obs.closed, + "MQTT-3.2.2-4 VIOLATION: brand-new client with no session state accepted CONNACK Session Present=1: connect result {:?}, network closed={}", + s.result, + obs.closed + ); +} + +#[tokio::test] +async fn mqtt_3_2_2_18_topic_alias_absent_client_must_not_send_alias() { + let mut s = start(opts("ta-absent"), false, &[]).await; + s.result.as_ref().unwrap(); + assert_no_alias_sent(&s.client, &mut s.stream, 1, "MQTT-3.2.2-18").await; +} + +#[tokio::test] +async fn mqtt_3_2_2_17_topic_alias_above_server_maximum() { + let mut s = start(opts("ta-max"), false, &p_u16(P_TOPIC_ALIAS_MAXIMUM, 2)).await; + s.result.as_ref().unwrap(); + assert_no_alias_sent(&s.client, &mut s.stream, 3, "MQTT-3.2.2-17").await; +} + +#[tokio::test] +async fn mqtt_3_2_2_14_retain_available_0_client_must_not_send_retain() { + let mut s = start(opts("ra0"), false, &p_u8(P_RETAIN_AVAILABLE, 0)).await; + s.result.as_ref().unwrap(); + let r = s.client.publish_retain("r/t", b"x".to_vec()).await; + if let Some(raw) = next_of_type(&mut s.stream, PUBLISH, Duration::from_secs(1)).await { + assert!( + !parse_publish(&raw).retain, + "MQTT-3.2.2-14 VIOLATION: client sent PUBLISH with RETAIN=1 after CONNACK Retain Available=0 (publish returned {r:?})" + ); + } +} + +#[tokio::test] +async fn mqtt_3_2_2_11_maximum_qos_live_publish() { + let mut s = start(opts("mq0"), false, &p_u8(P_MAXIMUM_QOS, 0)).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let h = + tokio::spawn(async move { c.publish_qos("q/t", b"x".to_vec(), QoS::ExactlyOnce).await }); + if let Some(raw) = next_of_type(&mut s.stream, PUBLISH, Duration::from_secs(1)).await { + assert_eq!( + parse_publish(&raw).qos, + 0, + "MQTT-3.2.2-11 VIOLATION: PUBLISH QoS exceeds server Maximum QoS 0" + ); + } + h.abort(); +} + +async fn queued_publish_then_reconnect( + conn2_props: &[u8], + options: PublishOptions, + id: &str, +) -> Option { + let o = reconnecting_opts(id) + .with_clean_start(false) + .with_session_expiry_interval(300); + let s = start(o, false, &[]).await; + s.result.as_ref().unwrap(); + let Setup { + client, + listener, + stream, + .. + } = s; + drop(stream); + wait_connected(&client, false).await; + let queued = client + .publish_with_options("queued/t", b"offline".to_vec(), options) + .await; + assert!( + queued.is_ok(), + "publish while offline should queue: {queued:?}" + ); + let (mut stream2, _) = accept_next(&listener, true, conn2_props).await; + let raw = next_of_type(&mut stream2, PUBLISH, Duration::from_secs(3)).await; + raw.as_ref().map(parse_publish).inspect(|p| { + assert_eq!(p.topic, b"queued/t"); + assert_eq!(p.payload, b"offline"); + }) +} + +#[tokio::test] +async fn mqtt_3_2_2_11_maximum_qos_applies_to_messages_queued_while_offline() { + let p = + queued_publish_then_reconnect(&p_u8(P_MAXIMUM_QOS, 1), qos(QoS::ExactlyOnce), "mq-queued") + .await; + let p = p.expect("queued message was never sent"); + assert!( + p.qos <= 1, + "MQTT-3.2.2-11 VIOLATION: message queued offline was flushed as QoS {} after CONNACK Maximum QoS=1", + p.qos + ); +} + +#[tokio::test] +async fn mqtt_3_2_2_14_retain_available_applies_to_messages_queued_while_offline() { + let options = PublishOptions { + qos: QoS::AtLeastOnce, + retain: true, + ..Default::default() + }; + let p = queued_publish_then_reconnect(&p_u8(P_RETAIN_AVAILABLE, 0), options, "ra-queued").await; + if let Some(p) = p { + assert!( + !p.retain, + "MQTT-3.2.2-14 VIOLATION: message queued offline was flushed with RETAIN=1 after CONNACK Retain Available=0" + ); + } +} + +#[tokio::test] +async fn crosscheck1_mqtt_3_2_2_18_topic_alias_maximum_not_stale_after_reconnect() { + let (client, _l, mut stream) = + reconnect_with(&p_u16(P_TOPIC_ALIAS_MAXIMUM, 10), &[], "stale-ta").await; + assert_no_alias_sent( + &client, + &mut stream, + 1, + "MQTT-3.2.2-18 (stale Topic Alias Maximum)", + ) + .await; +} + +#[tokio::test] +async fn crosscheck1_mqtt_3_2_2_11_maximum_qos_not_stale_after_reconnect() { + let (client, _l, mut stream) = reconnect_with(&p_u8(P_MAXIMUM_QOS, 1), &[], "stale-mq").await; + let c = client.clone(); + let h = + tokio::spawn(async move { c.publish_qos("q/t", b"x".to_vec(), QoS::ExactlyOnce).await }); + let raw = next_of_type(&mut stream, PUBLISH, Duration::from_secs(2)) + .await + .expect("no PUBLISH"); + assert_eq!( + parse_publish(&raw).qos, + 2, + "stale capability: conn 2 CONNACK omitted Maximum QoS (=2) but client still downgraded using conn 1 Maximum QoS=1" + ); + h.abort(); +} + +#[tokio::test] +async fn crosscheck1_mqtt_3_2_2_15_maximum_packet_size_not_stale_after_reconnect() { + let (client, _l, mut stream) = + reconnect_with(&p_u32(P_MAXIMUM_PACKET_SIZE, 64), &[], "stale-mps").await; + let r = client.publish("big/t", vec![b'x'; 500]).await; + assert!( + r.is_ok(), + "stale capability: conn 2 CONNACK omitted Maximum Packet Size but client still enforced conn 1 limit 64: {r:?}" + ); + assert!(next_of_type(&mut stream, PUBLISH, Duration::from_secs(1)) + .await + .is_some()); +} + +#[tokio::test] +async fn crosscheck1_mqtt_3_2_2_15_maximum_packet_size_newly_imposed_on_reconnect() { + let (client, _l, mut stream) = + reconnect_with(&[], &p_u32(P_MAXIMUM_PACKET_SIZE, 64), "new-mps").await; + let r = client.publish("big/t", vec![b'x'; 500]).await; + let sent = next_of_type(&mut stream, PUBLISH, Duration::from_secs(1)).await; + assert!( + sent.is_none() || sent.as_ref().is_some_and(|p| p.wire_len <= 64), + "MQTT-3.2.2-15 VIOLATION: client sent {}-byte PUBLISH after conn 2 CONNACK Maximum Packet Size=64 (publish returned {r:?})", + sent.map_or(0, |p| p.wire_len) + ); +} + +#[tokio::test] +async fn crosscheck1_retain_available_not_stale_after_reconnect() { + let (client, _l, mut stream) = + reconnect_with(&p_u8(P_RETAIN_AVAILABLE, 0), &[], "stale-ra").await; + let r = client.publish_retain("r/t", b"x".to_vec()).await; + let raw = next_of_type(&mut stream, PUBLISH, Duration::from_secs(1)).await; + assert!( + raw.as_ref().is_some_and(|p| parse_publish(p).retain), + "stale capability: conn 2 CONNACK omitted Retain Available (=1) but retained publish was not sent with RETAIN=1 ({r:?})" + ); +} + +#[tokio::test] +async fn crosscheck1_receive_maximum_not_stale_after_reconnect() { + let (client, _l, mut stream) = + reconnect_with(&p_u16(P_RECEIVE_MAXIMUM, 1), &[], "stale-rm").await; + let mut handles = Vec::new(); + for i in 0..3u8 { + let c = client.clone(); + handles.push(tokio::spawn(async move { + c.publish_qos("rm/t", vec![i], QoS::AtLeastOnce).await + })); + } + let mut seen = 0; + while next_of_type(&mut stream, PUBLISH, Duration::from_millis(700)) + .await + .is_some() + { + seen += 1; + } + for h in handles { + h.abort(); + } + assert_eq!( + seen, 3, + "stale capability: conn 2 CONNACK omitted Receive Maximum (=65535) but client only sent {seen} of 3 unacked QoS1 PUBLISHes" + ); +} + +#[tokio::test] +async fn crosscheck1_mqtt_3_2_2_11_maximum_qos_newly_imposed_on_reconnect() { + let (client, _l, mut stream) = reconnect_with(&[], &p_u8(P_MAXIMUM_QOS, 0), "new-mq").await; + let c = client.clone(); + let h = + tokio::spawn(async move { c.publish_qos("q/t", b"x".to_vec(), QoS::AtLeastOnce).await }); + if let Some(raw) = next_of_type(&mut stream, PUBLISH, Duration::from_secs(1)).await { + assert_eq!( + parse_publish(&raw).qos, + 0, + "MQTT-3.2.2-11 VIOLATION: after conn 2 CONNACK Maximum QoS=0 the client still sent QoS>0" + ); + } + h.abort(); +} + +fn rpi0_opts(id: &str) -> ConnectOptions { + let mut o = opts(id); + o.properties.request_problem_information = Some(false); + o +} + +#[tokio::test] +async fn crosscheck3_mqtt_3_1_2_29_request_problem_information_0_is_sent() { + let s = start(rpi0_opts("rpi-wire"), false, &[]).await; + let c = parse_connect(&s.connect); + assert!( + c.props.iter().any(|(id, v)| *id == P_REQUEST_PROBLEM_INFO && v == &[0]), + "MQTT-3.1.2-29 precondition BUG: ConnectProperties.request_problem_information=Some(false) was not put on the wire; CONNECT props {:?}", + c.props + ); +} + +#[tokio::test] +async fn crosscheck3_mqtt_3_1_2_29_reason_string_on_puback_with_rpi0_disconnects_0x82() { + let mut s = start(rpi0_opts("rpi-puback"), false, &[]).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let h = + tokio::spawn(async move { c.publish_qos("p/t", b"x".to_vec(), QoS::AtLeastOnce).await }); + let raw = next_of_type(&mut s.stream, PUBLISH, T).await.unwrap(); + let pid = parse_publish(&raw).packet_id.unwrap(); + s.stream + .write_all(&ack_with_props( + 0x40, + pid, + 0x00, + &p_str(P_REASON_STRING, b"why"), + )) + .await + .unwrap(); + let obs = observe_close(&mut s.stream, Duration::from_secs(2)).await; + let r = timeout(Duration::from_secs(1), h).await; + assert!( + obs.disconnect_rc == Some(0x82) && obs.closed, + "MQTT-3.1.2-29 (§3.1.2.11.7) VIOLATION: with Request Problem Information=0 configured, PUBACK carrying a Reason String did not trigger DISCONNECT 0x82 (disconnect rc {:?}, closed {}, publish result {r:?})", + obs.disconnect_rc, + obs.closed + ); +} + +#[tokio::test] +async fn crosscheck3_mqtt_3_1_2_29_user_property_on_suback_with_rpi0_disconnects_0x82() { + let mut s = start(rpi0_opts("rpi-suback"), false, &[]).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let h = tokio::spawn(async move { c.subscribe("s/t", |_| {}).await }); + let raw = next_of_type(&mut s.stream, SUBSCRIBE, T).await.unwrap(); + let pid = u16::from_be_bytes([raw.body[0], raw.body[1]]); + let mut body = pid.to_be_bytes().to_vec(); + body.extend(with_props(&p_pair(b"k", b"v"))); + body.push(0x00); + s.stream.write_all(&packet(0x90, &body)).await.unwrap(); + let obs = observe_close(&mut s.stream, Duration::from_secs(2)).await; + let r = timeout(Duration::from_secs(1), h).await; + assert!( + obs.disconnect_rc == Some(0x82) && obs.closed, + "MQTT-3.1.2-29 (§3.1.2.11.7) VIOLATION: with Request Problem Information=0 configured, SUBACK carrying a User Property did not trigger DISCONNECT 0x82 (disconnect rc {:?}, closed {}, subscribe result {r:?})", + obs.disconnect_rc, + obs.closed + ); +} + +type AuthFut<'a, R> = + std::pin::Pin> + Send + 'a>>; + +struct StaticAuth; + +impl AuthHandler for StaticAuth { + fn handle_challenge<'a>( + &'a self, + _auth_method: &'a str, + _challenge_data: Option<&'a [u8]>, + ) -> AuthFut<'a, AuthResponse> { + Box::pin(async move { Ok(AuthResponse::Continue(b"resp".to_vec())) }) + } + + fn initial_response<'a>(&'a self, _auth_method: &'a str) -> AuthFut<'a, Option>> { + Box::pin(async move { Ok(Some(b"init".to_vec())) }) + } +} + +fn auth_packet(rc: u8, extra_props: &[u8]) -> Vec { + let mut props = p_str(P_AUTH_METHOD, b"TEST"); + props.extend_from_slice(extra_props); + let mut body = vec![rc]; + body.extend(with_props(&props)); + packet(0xF0, &body) +} + +#[tokio::test] +async fn crosscheck3_mqtt_3_1_2_29_reason_string_on_reauth_auth_with_rpi0_disconnects_0x82() { + let mut o = rpi0_opts("rpi-auth"); + o.properties.authentication_method = Some("TEST".to_string()); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("mqtt://{}", listener.local_addr().unwrap()); + let client = MqttClient::with_options(o.clone()); + client.set_auth_handler(StaticAuth).await; + let c = client.clone(); + let h = tokio::spawn(async move { Box::pin(c.connect_with_options(&url, o)).await }); + let (mut stream, _) = timeout(T, listener.accept()).await.unwrap().unwrap(); + read_raw(&mut stream).await.unwrap(); + stream + .write_all(&connack(false, &p_str(P_AUTH_METHOD, b"TEST"))) + .await + .unwrap(); + timeout(T, h).await.unwrap().unwrap().unwrap(); + client.reauthenticate().await.unwrap(); + let auth = next_of_type(&mut stream, AUTH, T) + .await + .expect("no re-auth AUTH"); + assert_eq!(auth.body[0], 0x19); + stream + .write_all(&auth_packet(0x00, &p_str(P_REASON_STRING, b"why"))) + .await + .unwrap(); + let obs = observe_close(&mut stream, Duration::from_secs(2)).await; + assert!( + obs.disconnect_rc == Some(0x82) && obs.closed, + "MQTT-3.1.2-29 (§3.1.2.11.7) VIOLATION: with Request Problem Information=0 configured, re-auth AUTH carrying a Reason String did not trigger DISCONNECT 0x82 (disconnect rc {:?}, closed {})", + obs.disconnect_rc, + obs.closed + ); +} + +#[tokio::test] +async fn crosscheck4_mqtt_3_2_2_15_automatic_puback_respects_tiny_maximum_packet_size() { + let mut s = start(opts("tiny-mps"), false, &p_u32(P_MAXIMUM_PACKET_SIZE, 3)).await; + s.result.as_ref().unwrap(); + s.stream + .write_all(&publish_bytes(0x32, b"a", Some(7), &[], b"x")) + .await + .unwrap(); + let ack = next_of_type(&mut s.stream, PUBACK, Duration::from_secs(1)).await; + assert!( + ack.as_ref().is_none_or(|a| a.wire_len <= 3), + "MQTT-3.2.2-15 VIOLATION: automatic PUBACK of {} bytes sent despite server Maximum Packet Size=3", + ack.map_or(0, |a| a.wire_len) + ); +} + +#[tokio::test] +async fn crosscheck4_mqtt_3_2_2_15_automatic_pubrec_respects_tiny_maximum_packet_size() { + let mut s = start(opts("tiny-mps2"), false, &p_u32(P_MAXIMUM_PACKET_SIZE, 3)).await; + s.result.as_ref().unwrap(); + s.stream + .write_all(&publish_bytes(0x34, b"a", Some(8), &[], b"x")) + .await + .unwrap(); + let ack = next_of_type(&mut s.stream, PUBREC, Duration::from_secs(1)).await; + assert!( + ack.as_ref().is_none_or(|a| a.wire_len <= 3), + "MQTT-3.2.2-15 VIOLATION: automatic PUBREC of {} bytes sent despite server Maximum Packet Size=3", + ack.map_or(0, |a| a.wire_len) + ); +} + +#[tokio::test] +async fn crosscheck5_auth_data_without_auth_method_never_sent() { + let s = start( + opts("authdata").with_authentication_data(b"secret"), + false, + &[], + ) + .await; + let c = parse_connect(&s.connect); + assert!( + !has_prop(&c.props, P_AUTH_DATA) || has_prop(&c.props, P_AUTH_METHOD), + "§3.1.2.11.10 VIOLATION: CONNECT carries Authentication Data without Authentication Method: {:?}", + c.props + ); +} + +#[tokio::test] +async fn crosscheck6_mqtt_3_2_2_15_oversized_publish_rejected_and_packet_id_not_leaked() { + let mut s = start(opts("mps-pub"), false, &p_u32(P_MAXIMUM_PACKET_SIZE, 64)).await; + s.result.as_ref().unwrap(); + let big = s + .client + .publish_qos("p/t", vec![b'x'; 200], QoS::AtLeastOnce) + .await; + assert!( + matches!(big, Err(MqttError::PacketTooLarge { .. })), + "oversized publish not rejected: {big:?}" + ); + let c = s.client.clone(); + let h = + tokio::spawn(async move { c.publish_qos("p/t", b"x".to_vec(), QoS::AtLeastOnce).await }); + let raw = next_of_type(&mut s.stream, PUBLISH, T).await.unwrap(); + let p = parse_publish(&raw); + assert_eq!( + p.packet_id, + Some(1), + "packet id leaked by rejected oversized publish" + ); + s.stream + .write_all(&packet(0x40, &p.packet_id.unwrap().to_be_bytes())) + .await + .unwrap(); + timeout(T, h).await.unwrap().unwrap().unwrap(); +} + +#[tokio::test] +async fn crosscheck6_mqtt_3_2_2_15_oversized_subscribe_not_sent() { + let mut s = start(opts("mps-sub"), false, &p_u32(P_MAXIMUM_PACKET_SIZE, 64)).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let filter = format!("f/{}", "x".repeat(100)); + let h = tokio::spawn(async move { c.subscribe(filter, |_| {}).await }); + let sent = next_of_type(&mut s.stream, SUBSCRIBE, Duration::from_secs(1)).await; + h.abort(); + assert!( + sent.as_ref().is_none_or(|p| p.wire_len <= 64), + "MQTT-3.2.2-15 VIOLATION: client sent {}-byte SUBSCRIBE despite server Maximum Packet Size=64", + sent.map_or(0, |p| p.wire_len) + ); +} + +#[tokio::test] +async fn crosscheck6_mqtt_3_2_2_15_oversized_unsubscribe_not_sent() { + let mut s = start(opts("mps-unsub"), false, &p_u32(P_MAXIMUM_PACKET_SIZE, 64)).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let filter = format!("f/{}", "x".repeat(100)); + let h = tokio::spawn(async move { c.unsubscribe(filter).await }); + let sent = next_of_type(&mut s.stream, UNSUBSCRIBE, Duration::from_secs(1)).await; + h.abort(); + assert!( + sent.as_ref().is_none_or(|p| p.wire_len <= 64), + "MQTT-3.2.2-15 VIOLATION: client sent {}-byte UNSUBSCRIBE despite server Maximum Packet Size=64", + sent.map_or(0, |p| p.wire_len) + ); +} + +#[tokio::test] +async fn mqtt_2_2_1_3_packet_id_not_reused_while_in_flight_after_wrap() { + let mut s = start(opts("pid-wrap"), false, &[]).await; + s.result.as_ref().unwrap(); + let held = s.client.clone(); + let _held = tokio::spawn(async move { + held.publish_qos("w/held", b"h".to_vec(), QoS::AtLeastOnce) + .await + }); + let first = parse_publish(&next_of_type(&mut s.stream, PUBLISH, T).await.unwrap()); + let held_id = first.packet_id.unwrap(); + + let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::(); + let mut stream = s.stream; + let acker = tokio::spawn(async move { + while let Some(raw) = read_raw(&mut stream).await { + if raw.ptype() == PUBLISH { + let id = parse_publish(&raw).packet_id.unwrap(); + let _ = tx.send(id); + if id != held_id { + let _ = stream.write_all(&packet(0x40, &id.to_be_bytes())).await; + } + } else if raw.ptype() == PINGREQ { + let _ = stream.write_all(&[0xD0, 0x00]).await; + } + } + }); + + let mut reused = None; + for _ in 0..65535u32 { + let c = s.client.clone(); + let _ = timeout(T, c.publish_qos("w/t", b"x".to_vec(), QoS::AtLeastOnce)).await; + if let Some(id) = rx.recv().await { + if id == held_id { + reused = Some(id); + break; + } + } + } + acker.abort(); + assert!( + reused.is_none(), + "MQTT-2.2.1-3 VIOLATION: packet id {held_id} reassigned to a new PUBLISH while the original QoS1 PUBLISH with that id was still unacknowledged" + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_23_unacked_qos1_publish_resent_on_session_resume() { + let o = reconnecting_opts("resend1") + .with_clean_start(false) + .with_session_expiry_interval(300); + let mut s = start(o, false, &[]).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let _h = + tokio::spawn(async move { c.publish_qos("r/t", b"x".to_vec(), QoS::AtLeastOnce).await }); + let orig = parse_publish(&next_of_type(&mut s.stream, PUBLISH, T).await.unwrap()); + let Setup { + client, + listener, + stream, + .. + } = s; + drop(stream); + wait_connected(&client, false).await; + let (mut stream2, _) = accept_next(&listener, true, &[]).await; + let resent = next_of_type(&mut stream2, PUBLISH, Duration::from_secs(3)).await; + let resent = resent.map(|r| parse_publish(&r)); + assert!( + resent + .as_ref() + .is_some_and(|p| p.packet_id == orig.packet_id && p.dup), + "MQTT-3.1.2-23 / MQTT-4.4.0-1 VIOLATION: unacknowledged QoS1 PUBLISH id {:?} not resent (DUP=1, same id) after reconnect with Clean Start=0 and Session Present=1; got {resent:?}", + orig.packet_id + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_23_unacked_pubrel_resent_on_session_resume() { + let o = reconnecting_opts("resend2") + .with_clean_start(false) + .with_session_expiry_interval(300); + let mut s = start(o, false, &[]).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let _h = + tokio::spawn(async move { c.publish_qos("r/t", b"x".to_vec(), QoS::ExactlyOnce).await }); + let orig = parse_publish(&next_of_type(&mut s.stream, PUBLISH, T).await.unwrap()); + let pid = orig.packet_id.unwrap(); + s.stream + .write_all(&packet(0x50, &pid.to_be_bytes())) + .await + .unwrap(); + next_of_type(&mut s.stream, PUBREL, T) + .await + .expect("no PUBREL"); + let Setup { + client, + listener, + stream, + .. + } = s; + drop(stream); + wait_connected(&client, false).await; + let (mut stream2, _) = accept_next(&listener, true, &[]).await; + let resent = next_of_type(&mut stream2, PUBREL, Duration::from_secs(3)).await; + assert!( + resent + .as_ref() + .is_some_and(|r| u16::from_be_bytes([r.body[0], r.body[1]]) == pid), + "MQTT-3.1.2-23 / MQTT-4.4.0-1 VIOLATION: PUBREL for id {pid} not resent after reconnect with Clean Start=0 and Session Present=1" + ); +} + +#[tokio::test] +async fn mqtt_3_1_3_2_assigned_client_identifier_used_to_resume_session() { + let o = reconnecting_opts("") + .with_clean_start(false) + .with_session_expiry_interval(300); + let s = start(o, false, &p_str(P_ASSIGNED_CLIENT_ID, b"srv-assigned-1")).await; + s.result.as_ref().unwrap(); + let first = parse_connect(&s.connect); + let Setup { + client, + listener, + stream, + .. + } = s; + drop(stream); + wait_connected(&client, false).await; + let (_stream2, raw2) = accept_next(&listener, true, &[]).await; + let second = parse_connect(&raw2); + assert_eq!( + String::from_utf8_lossy(&second.client_id), + "srv-assigned-1", + "MQTT-3.1.3-2 VIOLATION: client holding Session State (Clean Start=0, SEI=300) under server-assigned id reconnected with ClientID {:?} (first CONNECT used {:?}), so the state it holds cannot be identified", + String::from_utf8_lossy(&second.client_id), + String::from_utf8_lossy(&first.client_id) + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_21_server_keep_alive_overrides_client_value() { + let mut s = start(opts("ska"), false, &p_u16(P_SERVER_KEEP_ALIVE, 1)).await; + s.result.as_ref().unwrap(); + let ping = next_of_type(&mut s.stream, PINGREQ, Duration::from_millis(1600)).await; + assert!( + ping.is_some(), + "MQTT-3.1.2-21 VIOLATION: Server Keep Alive=1 but no PINGREQ within 1.6s (client used its own 60s)" + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_20_pingreq_sent_when_idle() { + let o = opts("ka1").with_keep_alive(Duration::from_secs(1)); + let mut s = start(o, false, &[]).await; + s.result.as_ref().unwrap(); + let mut pings = 0; + for _ in 0..3 { + match next_packet(&mut s.stream, Duration::from_millis(1100)).await { + Next::Packet(r) if r.ptype() == PINGREQ => { + pings += 1; + s.stream.write_all(&[0xD0, 0x00]).await.unwrap(); + } + _ => break, + } + } + assert_eq!( + pings, 3, + "MQTT-3.1.2-20 VIOLATION: Keep Alive=1s idle client did not send a PINGREQ within each 1.1s window" + ); +} + +struct MalformedOutcome { + delivered: usize, + still_connected: bool, + obs: CloseObservation, +} + +async fn inject_after_subscribe(bytes: &[u8], id: &str) -> MalformedOutcome { + let mut s = start(opts(id), false, &[]).await; + s.result.as_ref().unwrap(); + let delivered = Arc::new(AtomicUsize::new(0)); + let d = delivered.clone(); + let c = s.client.clone(); + let sub = tokio::spawn(async move { + c.subscribe("#", move |_| { + d.fetch_add(1, Ordering::SeqCst); + }) + .await + }); + let raw = next_of_type(&mut s.stream, SUBSCRIBE, T).await.unwrap(); + let mut body = raw.body[0..2].to_vec(); + body.extend([0x00, 0x00]); + s.stream.write_all(&packet(0x90, &body)).await.unwrap(); + timeout(T, sub).await.unwrap().unwrap().unwrap(); + s.stream.write_all(bytes).await.unwrap(); + let obs = observe_close(&mut s.stream, Duration::from_secs(3)).await; + MalformedOutcome { + delivered: delivered.load(Ordering::SeqCst), + still_connected: s.client.is_connected().await, + obs, + } +} + +async fn assert_malformed_rejected(bytes: &[u8], id: &str, what: &str) { + let o = inject_after_subscribe(bytes, id).await; + assert!( + o.delivered == 0 && !o.still_connected, + "{id} VIOLATION: malformed {what} was accepted (delivered {} times, client still connected={})", + o.delivered, + o.still_connected + ); +} + +fn surrogate_topic_publish() -> Vec { + publish_bytes(0x30, &[b'a', 0xED, 0xA0, 0x80], None, &[], b"x") +} + +#[tokio::test] +async fn mqtt_1_5_4_1_surrogate_in_inbound_topic_is_malformed() { + assert_malformed_rejected( + &surrogate_topic_publish(), + "MQTT-1.5.4-1", + "PUBLISH topic with UTF-16 surrogate", + ) + .await; +} + +#[tokio::test] +async fn mqtt_1_5_4_2_null_in_inbound_topic_is_malformed() { + assert_malformed_rejected( + &publish_bytes(0x30, b"a\0b", None, &[], b"x"), + "MQTT-1.5.4-2", + "PUBLISH topic containing U+0000", + ) + .await; +} + +#[tokio::test] +async fn mqtt_1_5_7_1_invalid_utf8_user_property_pair_is_malformed() { + let mut prop = vec![P_USER_PROPERTY]; + prop.extend(mqtt_str(b"k")); + prop.extend(mqtt_str(&[0xFF, 0xFE])); + assert_malformed_rejected( + &publish_bytes(0x30, b"a", None, &prop, b"x"), + "MQTT-1.5.7-1", + "PUBLISH with invalid UTF-8 User Property value", + ) + .await; +} + +#[tokio::test] +async fn mqtt_2_1_3_1_reserved_flags_on_inbound_puback_is_malformed() { + let mut s = start(opts("flags-puback"), false, &[]).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let h = + tokio::spawn(async move { c.publish_qos("p/t", b"x".to_vec(), QoS::AtLeastOnce).await }); + let pid = parse_publish(&next_of_type(&mut s.stream, PUBLISH, T).await.unwrap()) + .packet_id + .unwrap(); + s.stream + .write_all(&packet(0x42, &pid.to_be_bytes())) + .await + .unwrap(); + let r = timeout(Duration::from_secs(2), h).await.unwrap().unwrap(); + assert!( + r.is_err() && !s.client.is_connected().await, + "MQTT-2.1.3-1 VIOLATION: PUBACK with reserved flag bits 0x2 accepted as a valid acknowledgement ({r:?})" + ); +} + +#[tokio::test] +async fn mqtt_2_1_3_1_reserved_flags_on_inbound_suback_is_malformed() { + let mut s = start(opts("flags-suback"), false, &[]).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let h = tokio::spawn(async move { c.subscribe("s/t", |_| {}).await }); + let raw = next_of_type(&mut s.stream, SUBSCRIBE, T).await.unwrap(); + let mut body = raw.body[0..2].to_vec(); + body.extend([0x00, 0x00]); + s.stream.write_all(&packet(0x92, &body)).await.unwrap(); + let r = timeout(Duration::from_secs(2), h).await.unwrap().unwrap(); + assert!( + r.is_err() && !s.client.is_connected().await, + "MQTT-2.1.3-1 VIOLATION: SUBACK with reserved flag bits 0x2 accepted as a valid acknowledgement ({r:?}, still connected={})", + s.client.is_connected().await + ); +} + +#[tokio::test] +async fn mqtt_4_13_2_1_network_connection_closed_after_malformed_packet() { + let o = inject_after_subscribe(&surrogate_topic_publish(), "close-after-malformed").await; + assert!( + o.obs.closed, + "MQTT-4.13.2-1 / §4.13.1 VIOLATION: after detecting a Malformed Packet (reason 0x81) the client marked itself disconnected (is_connected={}) but left the TCP connection open for 3s without sending DISCONNECT (rc {:?}); socket stays half-open until a reconnect or disconnect() runs", + o.still_connected, + o.obs.disconnect_rc + ); +} + +#[tokio::test] +async fn mqtt_3_2_2_1_connack_reserved_bits_rejected() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("mqtt://{}", listener.local_addr().unwrap()); + let o = opts("connack-flags"); + let client = MqttClient::with_options(o.clone()); + let c = client.clone(); + let h = tokio::spawn(async move { Box::pin(c.connect_with_options(&url, o)).await }); + let (mut stream, _) = timeout(T, listener.accept()).await.unwrap().unwrap(); + read_raw(&mut stream).await.unwrap(); + stream + .write_all(&[0x20, 0x03, 0x02, 0x00, 0x00]) + .await + .unwrap(); + let r = timeout(T, h).await.unwrap().unwrap(); + let obs = observe_close(&mut stream, Duration::from_secs(2)).await; + assert!( + r.is_err() && obs.closed, + "MQTT-3.2.2-1 VIOLATION: CONNACK with reserved flag bit 1 set accepted: {r:?}, closed {}", + obs.closed + ); +} + +#[tokio::test] +async fn mqtt_2_2_1_2_and_2_2_2_1_qos0_publish_has_no_packet_id_and_zero_property_length() { + let mut s = start(opts("qos0"), false, &[]).await; + s.result.as_ref().unwrap(); + s.client.publish("t", b"p".to_vec()).await.unwrap(); + let raw = next_of_type(&mut s.stream, PUBLISH, T).await.unwrap(); + assert_eq!( + raw.body, + vec![0x00, 0x01, b't', 0x00, b'p'], + "MQTT-2.2.1-2 / MQTT-2.2.2-1 VIOLATION: QoS0 PUBLISH body must be topic, property length 0, payload (no packet id)" + ); + assert_eq!(raw.flags(), 0, "MQTT-2.1.3-1: QoS0 PUBLISH flags"); +} + +#[tokio::test] +async fn mqtt_1_5_5_1_remaining_length_minimal_encoding() { + let mut s = start(opts("varint"), false, &[]).await; + s.result.as_ref().unwrap(); + for size in [100usize, 125, 126, 200, 16_380, 20_000] { + s.client.publish("v", vec![0u8; size]).await.unwrap(); + let raw = next_of_type(&mut s.stream, PUBLISH, T).await.unwrap(); + assert_eq!( + raw.remaining_len_bytes, + varint(raw.body.len()), + "MQTT-1.5.5-1 VIOLATION: non-minimal Remaining Length encoding" + ); + } +} + +#[tokio::test] +async fn mqtt_2_2_1_5_acks_echo_inbound_packet_id() { + let mut s = start(opts("ack-ids"), false, &[]).await; + s.result.as_ref().unwrap(); + s.stream + .write_all(&publish_bytes(0x32, b"a", Some(0x1234), &[], b"x")) + .await + .unwrap(); + let puback = next_of_type(&mut s.stream, PUBACK, T).await.unwrap(); + assert_eq!(puback.flags(), 0, "MQTT-2.1.3-1 PUBACK flags"); + assert_eq!(&puback.body[0..2], &[0x12, 0x34], "MQTT-2.2.1-5 PUBACK id"); + s.stream + .write_all(&publish_bytes(0x34, b"a", Some(0x4321), &[], b"x")) + .await + .unwrap(); + let pubrec = next_of_type(&mut s.stream, PUBREC, T).await.unwrap(); + assert_eq!(pubrec.flags(), 0, "MQTT-2.1.3-1 PUBREC flags"); + assert_eq!(&pubrec.body[0..2], &[0x43, 0x21], "MQTT-2.2.1-5 PUBREC id"); + s.stream + .write_all(&packet(0x62, &0x4321u16.to_be_bytes())) + .await + .unwrap(); + let pubcomp = next_of_type(&mut s.stream, PUBCOMP, T).await.unwrap(); + assert_eq!(pubcomp.flags(), 0, "MQTT-2.1.3-1 PUBCOMP flags"); + assert_eq!( + &pubcomp.body[0..2], + &[0x43, 0x21], + "MQTT-2.2.1-5 PUBCOMP id" + ); +} + +#[tokio::test] +async fn mqtt_2_1_3_1_outbound_reserved_flags_and_2_2_1_3_nonzero_ids() { + let mut s = start(opts("out-flags"), false, &[]).await; + s.result.as_ref().unwrap(); + let cl = s.client.clone(); + let task = tokio::spawn(async move { cl.subscribe("s/t", |_| {}).await }); + let sub = next_of_type(&mut s.stream, SUBSCRIBE, T).await.unwrap(); + assert_eq!(sub.flags(), 0x02, "MQTT-2.1.3-1 SUBSCRIBE flags"); + assert_ne!(&sub.body[0..2], &[0, 0], "MQTT-2.2.1-3 zero SUBSCRIBE id"); + let mut body = sub.body[0..2].to_vec(); + body.extend([0x00, 0x00]); + s.stream.write_all(&packet(0x90, &body)).await.unwrap(); + timeout(T, task).await.unwrap().unwrap().unwrap(); + + let cl = s.client.clone(); + let task = tokio::spawn(async move { cl.unsubscribe("s/t").await }); + let unsub = next_of_type(&mut s.stream, UNSUBSCRIBE, T).await.unwrap(); + assert_eq!(unsub.flags(), 0x02, "MQTT-2.1.3-1 UNSUBSCRIBE flags"); + assert_ne!( + &unsub.body[0..2], + &[0, 0], + "MQTT-2.2.1-3 zero UNSUBSCRIBE id" + ); + let mut body = unsub.body[0..2].to_vec(); + body.extend([0x00, 0x00]); + s.stream.write_all(&packet(0xB0, &body)).await.unwrap(); + timeout(T, task).await.unwrap().unwrap().unwrap(); + + let cl = s.client.clone(); + let task = + tokio::spawn(async move { cl.publish_qos("q2", b"x".to_vec(), QoS::ExactlyOnce).await }); + let publish = next_of_type(&mut s.stream, PUBLISH, T).await.unwrap(); + let pid = parse_publish(&publish).packet_id.unwrap(); + assert_ne!(pid, 0, "MQTT-2.2.1-3 zero PUBLISH id"); + s.stream + .write_all(&packet(0x50, &pid.to_be_bytes())) + .await + .unwrap(); + let rel = next_of_type(&mut s.stream, PUBREL, T).await.unwrap(); + assert_eq!(rel.flags(), 0x02, "MQTT-2.1.3-1 PUBREL flags"); + assert_eq!( + &rel.body[0..2], + &pid.to_be_bytes(), + "MQTT-2.2.1-5 PUBREL id" + ); + s.stream + .write_all(&packet(0x70, &pid.to_be_bytes())) + .await + .unwrap(); + timeout(T, task).await.unwrap().unwrap().unwrap(); + + s.client.disconnect().await.unwrap(); + let disc = next_of_type(&mut s.stream, DISCONNECT, T).await.unwrap(); + assert_eq!(disc.flags(), 0, "MQTT-2.1.3-1 DISCONNECT flags"); +} + +#[tokio::test] +async fn mqtt_3_1_connect_flags_payload_order_and_will() { + let will = mqtt5::WillMessage::new("will/t", b"bye".to_vec()).with_qos(QoS::AtLeastOnce); + let o = opts("cid-1") + .with_will(will) + .with_credentials("user", b"pass"); + let s = start(o, false, &[]).await; + let c = parse_connect(&s.connect); + assert_eq!(c.flags & 0x01, 0, "MQTT-3.1.2-3 reserved CONNECT flag"); + assert_eq!(c.flags & 0x04, 0x04, "will flag"); + assert!( + c.will_props.is_some(), + "MQTT-3.1.2-9 will properties present" + ); + assert_eq!( + c.will_topic.as_deref(), + Some(&b"will/t"[..]), + "MQTT-3.1.2-9 will topic" + ); + assert_eq!( + c.will_payload.as_deref(), + Some(&b"bye"[..]), + "MQTT-3.1.2-9 will payload" + ); + assert_eq!( + c.username.as_deref(), + Some(&b"user"[..]), + "MQTT-3.1.2-17 / 3.1.3-1 username" + ); + assert_eq!( + c.password.as_deref(), + Some(&b"pass"[..]), + "MQTT-3.1.2-19 / 3.1.3-1 password" + ); + assert_eq!(c.client_id, b"cid-1", "MQTT-3.1.3-1 client id first"); + assert_eq!( + c.trailing, 0, + "MQTT-3.1.3-1 no trailing bytes after declared fields" + ); + assert_eq!(c.keep_alive, 60); + + let s2 = start(opts("cid-2"), false, &[]).await; + let c2 = parse_connect(&s2.connect); + assert_eq!( + c2.flags & 0x3C, + 0, + "MQTT-3.1.2-13 will QoS/retain must be 0 without will" + ); + assert_eq!(c2.flags & 0xC0, 0, "MQTT-3.1.2-16/18 no user/pass flags"); + assert_eq!(c2.trailing, 0, "MQTT-3.1.2-16/18 no user/pass present"); + assert!( + c2.props.is_empty(), + "MQTT-2.2.2-1: CONNECT without properties" + ); + assert_eq!(s2.connect.body[10], 0x00, "MQTT-2.2.2-1: Property Length 0"); +} + +#[tokio::test] +async fn mqtt_1_5_4_2_null_character_never_encoded() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("mqtt://{}", listener.local_addr().unwrap()); + let options = opts("nul").with_credentials("us\0er", b"p"); + let client = MqttClient::with_options(options.clone()); + let cl = client.clone(); + let task = tokio::spawn(async move { Box::pin(cl.connect_with_options(&url, options)).await }); + let (mut stream, _) = timeout(T, listener.accept()).await.unwrap().unwrap(); + let got = timeout(Duration::from_secs(2), read_raw(&mut stream)).await; + let outcome = timeout(T, task).await; + assert!( + !matches!(got, Ok(Some(ref raw)) if raw.body.windows(5).any(|w| w == b"us\0er")), + "MQTT-1.5.4-2 VIOLATION: CONNECT User Name carries U+0000 (connect {outcome:?})" + ); + + let mut s = start(opts("nul-topic"), false, &[]).await; + s.result.as_ref().unwrap(); + let outcome = s.client.publish("a\0b", b"x".to_vec()).await; + let sent = next_of_type(&mut s.stream, PUBLISH, Duration::from_millis(500)).await; + assert!( + sent.is_none(), + "MQTT-1.5.4-2 VIOLATION: PUBLISH topic with U+0000 sent ({outcome:?})" + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_30_only_auth_before_connack_when_auth_method_set() { + let mut o = opts("auth-seq"); + o.properties.authentication_method = Some("TEST".to_string()); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("mqtt://{}", listener.local_addr().unwrap()); + let client = MqttClient::with_options(o.clone()); + client.set_auth_handler(StaticAuth).await; + let c = client.clone(); + let h = tokio::spawn(async move { Box::pin(c.connect_with_options(&url, o)).await }); + let (mut stream, _) = timeout(T, listener.accept()).await.unwrap().unwrap(); + let connect = read_raw(&mut stream).await.unwrap(); + let cp = parse_connect(&connect); + assert!(has_prop(&cp.props, P_AUTH_METHOD)); + let spam = client.clone(); + let spammer = tokio::spawn(async move { + for _ in 0..50 { + let _ = spam.publish("x", b"y".to_vec()).await; + let _ = spam.subscribe("z", |_| {}).await; + tokio::time::sleep(Duration::from_millis(10)).await; + } + }); + stream.write_all(&auth_packet(0x18, &[])).await.unwrap(); + let mut before_connack = Vec::new(); + while let Next::Packet(r) = next_packet(&mut stream, Duration::from_millis(500)).await { + before_connack.push(r.ptype()); + } + spammer.abort(); + stream + .write_all(&connack(false, &p_str(P_AUTH_METHOD, b"TEST"))) + .await + .unwrap(); + let _ = timeout(T, h).await; + assert!( + before_connack.iter().all(|t| *t == AUTH || *t == DISCONNECT), + "MQTT-3.1.2-30 VIOLATION: packets other than AUTH/DISCONNECT sent before CONNACK: {before_connack:?}" + ); + assert!( + before_connack.contains(&AUTH), + "client did not answer AUTH challenge" + ); +} + +#[tokio::test] +async fn json_mqtt_3_2_2_21_wildcard_subscription_available_0() { + let mut s = start(opts("wsa0"), false, &p_u8(P_WILDCARD_AVAILABLE, 0)).await; + s.result.as_ref().unwrap(); + let c = s.client.clone(); + let h = tokio::spawn(async move { c.subscribe("a/#", |_| {}).await }); + let sent = next_of_type(&mut s.stream, SUBSCRIBE, Duration::from_secs(1)).await; + h.abort(); + assert!( + sent.is_none(), + "manifest MQTT-3.2.2-21 text (Wildcard Subscription Available=0) VIOLATION: client sent SUBSCRIBE with wildcard filter" + ); +} + +#[tokio::test] +async fn resume_existing_session_opt_in_accepts_broker_held_session() { + let o = opts("sp-opt-in") + .with_clean_start(false) + .with_session_expiry_interval(300) + .with_resume_existing_session(true); + let mut s = start(o, true, &[]).await; + let obs = observe_close(&mut s.stream, Duration::from_millis(500)).await; + assert!( + s.result.as_ref().is_ok_and(|r| r.session_present) && !obs.closed, + "resume_existing_session(true) must accept Session Present=1 for a fresh Clean Start=0 client: connect result {:?}, network closed={}", + s.result, + obs.closed + ); +} + +#[tokio::test] +async fn resume_existing_session_opt_in_does_not_cover_clean_start_1() { + let mut s = start( + opts("sp-opt-in-clean").with_resume_existing_session(true), + true, + &[], + ) + .await; + let obs = observe_close(&mut s.stream, Duration::from_secs(2)).await; + assert!( + s.result.is_err() && obs.closed, + "MQTT-3.2.2-4 VIOLATION: Clean Start=1 accepted Session Present=1 despite the opt-in: connect result {:?}, network closed={}", + s.result, + obs.closed + ); +} + +#[tokio::test] +async fn topic_alias_not_carried_into_resumed_session_replay() { + let o = reconnecting_opts("alias-replay") + .with_clean_start(false) + .with_session_expiry_interval(300); + let mut s = start(o, false, &p_u16(P_TOPIC_ALIAS_MAXIMUM, 5)).await; + s.result.as_ref().unwrap(); + s.client + .publish_with_options("alias/t", b"map".to_vec(), with_alias(1)) + .await + .unwrap(); + let mapped = parse_publish(&next_of_type(&mut s.stream, PUBLISH, T).await.unwrap()); + assert_eq!(mapped.topic, b"alias/t"); + let c = s.client.clone(); + let _h = tokio::spawn(async move { + let options = PublishOptions { + qos: QoS::AtLeastOnce, + ..with_alias(1) + }; + c.publish_with_options("", b"aliased".to_vec(), options) + .await + }); + let orig = parse_publish(&next_of_type(&mut s.stream, PUBLISH, T).await.unwrap()); + assert!(orig.topic.is_empty(), "setup: live PUBLISH uses the alias"); + let Setup { + client, + listener, + stream, + .. + } = s; + drop(stream); + wait_connected(&client, false).await; + let (mut stream2, _) = accept_next(&listener, true, &p_u16(P_TOPIC_ALIAS_MAXIMUM, 5)).await; + let resent = parse_publish( + &next_of_type(&mut stream2, PUBLISH, Duration::from_secs(3)) + .await + .expect("unacked PUBLISH not resent"), + ); + assert_eq!(resent.packet_id, orig.packet_id); + assert_eq!( + resent.topic, b"alias/t", + "replayed PUBLISH must carry the full Topic Name; the alias mapping belongs to the previous connection" + ); + assert!( + !has_prop(&resent.props, P_TOPIC_ALIAS), + "MQTT-3.3.2-7 VIOLATION: replayed PUBLISH reused a Topic Alias never mapped on the new connection" + ); +} + +#[tokio::test] +async fn topic_alias_not_carried_into_offline_queue_flush() { + let options = PublishOptions { + qos: QoS::AtLeastOnce, + ..with_alias(1) + }; + let p = + queued_publish_then_reconnect(&p_u16(P_TOPIC_ALIAS_MAXIMUM, 5), options, "alias-queued") + .await + .expect("queued message was never sent"); + assert!( + !has_prop(&p.props, P_TOPIC_ALIAS), + "MQTT-3.3.2-7 VIOLATION: queued PUBLISH carried a Topic Alias created before the connection it was flushed on" + ); +} diff --git a/crates/mqtt5/tests/conf_client_b.rs b/crates/mqtt5/tests/conf_client_b.rs new file mode 100644 index 00000000..7dd3691d --- /dev/null +++ b/crates/mqtt5/tests/conf_client_b.rs @@ -0,0 +1,1509 @@ +use bytes::BytesMut; +use mqtt5::packet::connack::ConnAckPacket; +use mqtt5::packet::connect::ConnectPacket; +use mqtt5::packet::publish::PublishPacket; +use mqtt5::packet::{FixedHeader, MqttPacket, Packet, PacketType}; +use mqtt5::protocol::v5::reason_codes::ReasonCode; +use mqtt5::session::TopicAliasManager; +use mqtt5::{ + AckToken, ConnectOptions, Message, MqttClient, MqttError, PublishOptions, PublishProperties, + QoS, SubscribeOptions, +}; +use std::sync::{Arc, Mutex}; +use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; +use tokio::sync::mpsc; + +const CONNECT: u8 = 1; +const PUBLISH: u8 = 3; +const PUBACK: u8 = 4; +const PUBREC: u8 = 5; +const PUBREL: u8 = 6; +const PUBCOMP: u8 = 7; +const SUBSCRIBE: u8 = 8; +const UNSUBSCRIBE: u8 = 10; +const PINGREQ: u8 = 12; +const DISCONNECT: u8 = 14; + +const SHORT: Duration = Duration::from_millis(600); +const MEDIUM: Duration = Duration::from_secs(3); + +struct Frame { + first: u8, + body: Vec, +} + +impl Frame { + fn kind(&self) -> u8 { + self.first >> 4 + } + + fn dup(&self) -> bool { + self.first & 0x08 != 0 + } + + fn packet_id(&self) -> u16 { + if self.kind() == PUBLISH { + return self + .publish() + .packet_id + .expect("QoS > 0 PUBLISH carries a packet id"); + } + u16::from_be_bytes([self.body[0], self.body[1]]) + } + + fn reason_code(&self) -> u8 { + match self.kind() { + DISCONNECT => self.body.first().copied().unwrap_or(0), + _ => self.body.get(2).copied().unwrap_or(0), + } + } + + fn packet(&self) -> Packet { + let packet_type = PacketType::from_u8(self.kind()).expect("known packet type"); + let header = FixedHeader::new( + packet_type, + self.first & 0x0F, + u32::try_from(self.body.len()).expect("body fits u32"), + ); + let mut buf = BytesMut::from(&self.body[..]); + Packet::decode_from_body_with_version(packet_type, &header, &mut buf, 5) + .expect("client sent a decodable packet") + } + + fn publish(&self) -> PublishPacket { + match self.packet() { + Packet::Publish(publish) => publish, + other => panic!("expected PUBLISH, got {}", other.packet_type_name()), + } + } +} + +enum Next { + Frame(Frame), + Closed, + Timeout, +} + +async fn read_frame(stream: &mut TcpStream) -> Option { + let mut first = [0u8; 1]; + stream.read_exact(&mut first).await.ok()?; + let mut multiplier = 1usize; + let mut remaining = 0usize; + loop { + let mut byte = [0u8; 1]; + stream.read_exact(&mut byte).await.ok()?; + remaining += usize::from(byte[0] & 0x7F) * multiplier; + if byte[0] & 0x80 == 0 { + break; + } + multiplier *= 128; + } + let mut body = vec![0u8; remaining]; + stream.read_exact(&mut body).await.ok()?; + Some(Frame { + first: first[0], + body, + }) +} + +async fn next_frame(stream: &mut TcpStream, wait: Duration) -> Next { + let deadline = tokio::time::Instant::now() + wait; + loop { + match tokio::time::timeout_at(deadline, read_frame(stream)).await { + Err(_) => return Next::Timeout, + Ok(None) => return Next::Closed, + Ok(Some(frame)) if frame.kind() == PINGREQ => { + let _ = stream.write_all(&[0xD0, 0x00]).await; + } + Ok(Some(frame)) => return Next::Frame(frame), + } + } +} + +async fn expect_kind(stream: &mut TcpStream, kind: u8, wait: Duration) -> Frame { + match next_frame(stream, wait).await { + Next::Frame(frame) if frame.kind() == kind => frame, + Next::Frame(frame) => panic!("expected packet type {kind}, got {}", frame.kind()), + Next::Closed => panic!("connection closed while waiting for packet type {kind}"), + Next::Timeout => panic!("timed out waiting for packet type {kind}"), + } +} + +async fn collect_frames(stream: &mut TcpStream, wait: Duration) -> (Vec, bool) { + let deadline = tokio::time::Instant::now() + wait; + let mut frames = Vec::new(); + loop { + let left = deadline.saturating_duration_since(tokio::time::Instant::now()); + match next_frame(stream, left).await { + Next::Frame(frame) => frames.push(frame), + Next::Closed => return (frames, true), + Next::Timeout => return (frames, false), + } + } +} + +#[derive(Debug, PartialEq, Eq)] +enum Termination { + Disconnect(u8), + Closed, + StillOpen, +} + +async fn termination(stream: &mut TcpStream, wait: Duration) -> Termination { + let deadline = tokio::time::Instant::now() + wait; + loop { + let left = deadline.saturating_duration_since(tokio::time::Instant::now()); + match next_frame(stream, left).await { + Next::Frame(frame) if frame.kind() == DISCONNECT => { + return Termination::Disconnect(frame.reason_code()) + } + Next::Frame(_) => {} + Next::Closed => return Termination::Closed, + Next::Timeout => return Termination::StillOpen, + } + } +} + +fn encode_varint(mut value: usize, out: &mut Vec) { + loop { + let mut byte = u8::try_from(value % 128).expect("remainder fits u8"); + value /= 128; + if value > 0 { + byte |= 0x80; + } + out.push(byte); + if value == 0 { + break; + } + } +} + +fn frame_bytes(first: u8, body: &[u8]) -> Vec { + let mut out = vec![first]; + encode_varint(body.len(), &mut out); + out.extend_from_slice(body); + out +} + +fn raw_publish( + qos: u8, + dup: bool, + topic: &str, + packet_id: Option, + props: &[u8], + payload: &[u8], +) -> Vec { + let mut body = Vec::new(); + body.extend_from_slice( + &u16::try_from(topic.len()) + .expect("topic fits u16") + .to_be_bytes(), + ); + body.extend_from_slice(topic.as_bytes()); + if let Some(id) = packet_id { + body.extend_from_slice(&id.to_be_bytes()); + } + encode_varint(props.len(), &mut body); + body.extend_from_slice(props); + body.extend_from_slice(payload); + frame_bytes(0x30 | (u8::from(dup) << 3) | (qos << 1), &body) +} + +fn topic_alias_prop(alias: u16) -> Vec { + let mut out = vec![0x23]; + out.extend_from_slice(&alias.to_be_bytes()); + out +} + +fn subscription_id_prop(id: usize) -> Vec { + let mut out = vec![0x0B]; + encode_varint(id, &mut out); + out +} + +fn ack_bytes(first: u8, packet_id: u16) -> Vec { + frame_bytes(first, &packet_id.to_be_bytes()) +} + +fn suback_bytes(packet_id: u16, code: u8) -> Vec { + let mut body = packet_id.to_be_bytes().to_vec(); + body.push(0); + body.push(code); + frame_bytes(0x90, &body) +} + +fn unsuback_bytes(packet_id: u16) -> Vec { + let mut body = packet_id.to_be_bytes().to_vec(); + body.push(0); + body.push(0); + frame_bytes(0xB0, &body) +} + +fn connack(session_present: bool, configure: impl FnOnce(&mut ConnAckPacket)) -> Vec { + let mut packet = ConnAckPacket::new(session_present, ReasonCode::Success); + configure(&mut packet); + let mut encoded = Vec::new(); + packet.encode(&mut encoded).expect("CONNACK encodes"); + encoded +} + +fn plain_connack() -> Vec { + connack(false, |_| {}) +} + +async fn bind() -> (TcpListener, String) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let url = format!("mqtt://{}", listener.local_addr().expect("local addr")); + (listener, url) +} + +async fn accept_session( + listener: &TcpListener, + connack_bytes: &[u8], +) -> (TcpStream, ConnectPacket) { + let (mut stream, _) = tokio::time::timeout(Duration::from_secs(10), listener.accept()) + .await + .expect("client connected within 10s") + .expect("accept"); + let frame = expect_kind(&mut stream, CONNECT, MEDIUM).await; + let connect = match frame.packet() { + Packet::Connect(connect) => *connect, + other => panic!("expected CONNECT, got {}", other.packet_type_name()), + }; + stream + .write_all(connack_bytes) + .await + .expect("write CONNACK"); + (stream, connect) +} + +fn base_options(client_id: &str) -> ConnectOptions { + ConnectOptions::new(client_id).with_automatic_reconnect(false) +} + +fn resumable_options(client_id: &str) -> ConnectOptions { + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_session_expiry_interval(300) + .with_automatic_reconnect(true) + .with_reconnect_delay(Duration::from_millis(100), Duration::from_millis(200)) +} + +async fn connect_client( + options: ConnectOptions, + connack_bytes: &[u8], +) -> (MqttClient, TcpStream, TcpListener) { + let (listener, url) = bind().await; + let client = MqttClient::with_options(options); + let connecting = client.clone(); + let task = tokio::spawn(async move { connecting.connect(&url).await }); + let (stream, _) = accept_session(&listener, connack_bytes).await; + task.await + .expect("connect task") + .expect("client connects to fake broker"); + (client, stream, listener) +} + +async fn subscribe( + client: &MqttClient, + stream: &mut TcpStream, + filter: &str, + options: SubscribeOptions, +) -> mpsc::UnboundedReceiver { + let (tx, rx) = mpsc::unbounded_channel(); + let subscriber = client.clone(); + let filter = filter.to_string(); + let granted = options.qos as u8; + let task = tokio::spawn(async move { + subscriber + .subscribe_with_options(filter, options, move |message| { + let _ = tx.send(message); + }) + .await + }); + let frame = expect_kind(stream, SUBSCRIBE, MEDIUM).await; + stream + .write_all(&suback_bytes(frame.packet_id(), granted)) + .await + .expect("write SUBACK"); + task.await.expect("subscribe task").expect("subscribe"); + rx +} + +fn qos_options(qos: QoS) -> SubscribeOptions { + SubscribeOptions { + qos, + ..Default::default() + } +} + +async fn drain(rx: &mut mpsc::UnboundedReceiver, wait: Duration) -> Vec { + let mut messages = Vec::new(); + while let Ok(Some(message)) = tokio::time::timeout(wait, rx.recv()).await { + messages.push(message); + } + messages +} + +async fn wait_until_disconnected(client: &MqttClient) { + for _ in 0..100 { + if !client.is_connected().await { + return; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + panic!("client did not observe the dropped connection"); +} + +async fn published_frame_for( + options: PublishOptions, + topic: &str, +) -> (std::result::Result<(), MqttError>, Option) { + published_frame_with_connack(options, topic, &plain_connack()).await +} + +async fn published_frame_with_connack( + options: PublishOptions, + topic: &str, + connack_bytes: &[u8], +) -> (std::result::Result<(), MqttError>, Option) { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-pub"), connack_bytes).await; + let result = client + .publish_with_options(topic.to_string(), b"x".to_vec(), options) + .await + .map(|_| ()); + let frame = match next_frame(&mut stream, SHORT).await { + Next::Frame(frame) if frame.kind() == PUBLISH => Some(frame), + _ => None, + }; + (result, frame) +} + +fn qos0_with(properties: PublishProperties) -> PublishOptions { + PublishOptions { + qos: QoS::AtMostOnce, + properties, + ..Default::default() + } +} + +#[tokio::test] +async fn mqtt_3_3_2_2_client_publish_topic_must_not_contain_wildcards() { + let (result, frame) = + published_frame_for(qos0_with(PublishProperties::default()), "a/+/b").await; + if let Some(frame) = frame { + panic!( + "[MQTT-3.3.2-2] client sent PUBLISH with wildcard Topic Name {:?} (publish returned {result:?})", + frame.publish().topic_name + ); + } +} + +#[tokio::test] +async fn mqtt_3_3_2_1_prose_zero_length_topic_requires_topic_alias() { + let (result, frame) = published_frame_for(qos0_with(PublishProperties::default()), "").await; + if let Some(frame) = frame { + let publish = frame.publish(); + assert!( + publish.topic_alias().is_some(), + "[MQTT-3.3.2-1 / §3.3.2.3.4] client sent PUBLISH with zero-length Topic Name and no Topic Alias (publish returned {result:?})" + ); + } +} + +#[tokio::test] +async fn mqtt_3_3_2_8_client_must_not_send_topic_alias_zero() { + let properties = PublishProperties { + topic_alias: Some(0), + ..Default::default() + }; + let (result, frame) = published_frame_for(qos0_with(properties), "t/alias").await; + if let Some(frame) = frame { + assert_ne!( + frame.publish().topic_alias(), + Some(0), + "[MQTT-3.3.2-8] client sent PUBLISH with Topic Alias 0 (publish returned {result:?})" + ); + } +} + +#[tokio::test] +async fn mqtt_3_3_2_9_client_must_not_exceed_server_topic_alias_maximum() { + let properties = PublishProperties { + topic_alias: Some(3), + ..Default::default() + }; + let connack_bytes = connack(false, |c| c.properties.set_topic_alias_maximum(2)); + let (result, frame) = + published_frame_with_connack(qos0_with(properties), "t/alias", &connack_bytes).await; + if let Some(frame) = frame { + let alias = frame.publish().topic_alias(); + assert!( + alias.is_none_or(|a| a <= 2), + "[MQTT-3.3.2-9] client sent Topic Alias {alias:?} above the server Topic Alias Maximum 2 (publish returned {result:?})" + ); + } +} + +#[tokio::test] +async fn mqtt_3_3_4_6_client_publish_must_not_contain_subscription_identifier() { + let properties = PublishProperties { + subscription_identifiers: vec![7], + ..Default::default() + }; + let (result, frame) = published_frame_for(qos0_with(properties), "t/subid").await; + if let Some(frame) = frame { + let ids = frame + .publish() + .properties + .get_all(mqtt5::PropertyId::SubscriptionIdentifier) + .map_or(0, <[mqtt5::PropertyValue]>::len); + assert_eq!( + ids, 0, + "[MQTT-3.3.4-6] client sent PUBLISH containing a Subscription Identifier (publish returned {result:?})" + ); + } +} + +#[tokio::test] +async fn mqtt_3_3_2_14_response_topic_must_not_contain_wildcards() { + let properties = PublishProperties { + response_topic: Some("reply/#".to_string()), + ..Default::default() + }; + let (result, frame) = published_frame_for(qos0_with(properties), "t/req").await; + if let Some(frame) = frame { + panic!( + "[MQTT-3.3.2-14] client sent PUBLISH with wildcard Response Topic {:?} (publish returned {result:?})", + frame.publish().properties.get(mqtt5::PropertyId::ResponseTopic) + ); + } +} + +#[tokio::test] +async fn mqtt_3_3_2_1_13_19_utf8_string_fields_reject_nul() { + let cases = [ + ( + "t/\0nul", + PublishProperties::default(), + "MQTT-3.3.2-1 Topic Name", + ), + ( + "t/ok", + PublishProperties { + response_topic: Some("r/\0".to_string()), + ..Default::default() + }, + "MQTT-3.3.2-13 Response Topic", + ), + ( + "t/ok", + PublishProperties { + content_type: Some("text/\0".to_string()), + ..Default::default() + }, + "MQTT-3.3.2-19 Content Type", + ), + ]; + for (topic, properties, id) in cases { + let (result, frame) = published_frame_for(qos0_with(properties), topic).await; + assert!( + frame.is_none(), + "[{id}] client put a U+0000 string on the wire" + ); + assert!(result.is_err(), "[{id}] publish with U+0000 must fail"); + } +} + +#[tokio::test] +async fn mqtt_3_8_2_1_2_prose_subscribe_subscription_identifier_zero() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-subid0"), &plain_connack()).await; + let subscriber = client.clone(); + let task = tokio::spawn(async move { + subscriber + .subscribe_with_options( + "t/x", + SubscribeOptions { + subscription_identifier: Some(0), + ..Default::default() + }, + |_| {}, + ) + .await + }); + let sent = match next_frame(&mut stream, SHORT).await { + Next::Frame(frame) if frame.kind() == SUBSCRIBE => match frame.packet() { + Packet::Subscribe(sub) => sub.properties.get_subscription_identifier(), + _ => None, + }, + _ => None, + }; + task.abort(); + assert_ne!( + sent, + Some(0), + "[§3.8.2.1.2] client sent SUBSCRIBE with Subscription Identifier 0 (Protocol Error)" + ); +} + +#[tokio::test] +async fn mqtt_3_8_1_1_3_10_1_1_3_6_1_1_reserved_flags_and_reason_codes_on_wire() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-flags"), &plain_connack()).await; + + let subscriber = client.clone(); + let sub_task = tokio::spawn(async move { subscriber.subscribe("t/flags", |_| {}).await }); + let sub = expect_kind(&mut stream, SUBSCRIBE, MEDIUM).await; + assert_eq!( + sub.first, 0x82, + "[MQTT-3.8.1-1] SUBSCRIBE fixed header flags" + ); + match sub.packet() { + Packet::Subscribe(packet) => assert!( + !packet.filters.is_empty(), + "[MQTT-3.8.3-2] SUBSCRIBE must carry at least one filter" + ), + _ => unreachable!(), + } + stream + .write_all(&suback_bytes(sub.packet_id(), 0)) + .await + .unwrap(); + sub_task.await.unwrap().expect("subscribe"); + + client.publish("t/q0", b"x".to_vec()).await.expect("QoS0"); + let q0 = expect_kind(&mut stream, PUBLISH, MEDIUM).await; + assert!(!q0.dup(), "[MQTT-3.3.1-2] QoS0 PUBLISH must have DUP=0"); + + let publisher = client.clone(); + let q2_task = tokio::spawn(async move { publisher.publish_qos2("t/q2", b"x".to_vec()).await }); + let q2 = expect_kind(&mut stream, PUBLISH, MEDIUM).await; + assert!( + !q2.dup(), + "[MQTT-4.3.3-2] first QoS2 transmission must have DUP=0" + ); + stream + .write_all(&ack_bytes(0x50, q2.packet_id())) + .await + .unwrap(); + let rel = expect_kind(&mut stream, PUBREL, MEDIUM).await; + assert_eq!(rel.first, 0x62, "[MQTT-3.6.1-1] PUBREL fixed header flags"); + assert!( + matches!(rel.reason_code(), 0x00 | 0x92), + "[MQTT-3.6.2-1] PUBREL reason code 0x{:02X}", + rel.reason_code() + ); + assert_eq!(rel.packet_id(), q2.packet_id()); + stream + .write_all(&ack_bytes(0x70, q2.packet_id())) + .await + .unwrap(); + q2_task.await.unwrap().expect("QoS2 publish completes"); + + let unsubscriber = client.clone(); + let unsub_task = tokio::spawn(async move { unsubscriber.unsubscribe("t/flags").await }); + let unsub = expect_kind(&mut stream, UNSUBSCRIBE, MEDIUM).await; + assert_eq!( + unsub.first, 0xA2, + "[MQTT-3.10.1-1] UNSUBSCRIBE fixed header flags" + ); + match unsub.packet() { + Packet::Unsubscribe(packet) => assert!( + !packet.filters.is_empty(), + "[MQTT-3.10.3-2] UNSUBSCRIBE must carry at least one filter" + ), + _ => unreachable!(), + } + stream + .write_all(&unsuback_bytes(unsub.packet_id())) + .await + .unwrap(); + unsub_task.await.unwrap().expect("unsubscribe"); +} + +#[tokio::test] +async fn mqtt_3_8_4_2_suback_with_foreign_packet_id_is_not_accepted() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-suback"), &plain_connack()).await; + let subscriber = client.clone(); + let task = tokio::spawn(async move { subscriber.subscribe("t/s", |_| {}).await }); + let sub = expect_kind(&mut stream, SUBSCRIBE, MEDIUM).await; + let foreign = sub.packet_id().wrapping_add(100); + stream.write_all(&suback_bytes(foreign, 0)).await.unwrap(); + tokio::time::sleep(SHORT).await; + assert!( + !task.is_finished(), + "[MQTT-3.8.4-2] subscribe completed on a SUBACK carrying a different Packet Identifier" + ); + stream + .write_all(&suback_bytes(sub.packet_id(), 0)) + .await + .unwrap(); + tokio::time::timeout(MEDIUM, task) + .await + .expect("[MQTT-3.8.4-2] matching SUBACK completes subscribe") + .unwrap() + .expect("subscribe"); +} + +#[tokio::test] +async fn mqtt_3_3_4_1_3_4_2_3_5_2_3_7_2_client_acks_inbound_by_qos() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-acks"), &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/in", qos_options(QoS::ExactlyOnce)).await; + + stream + .write_all(&raw_publish(1, false, "t/in", Some(11), &[], b"q1")) + .await + .unwrap(); + let puback = expect_kind(&mut stream, PUBACK, MEDIUM).await; + assert_eq!(puback.packet_id(), 11, "[MQTT-3.3.4-1] PUBACK id"); + assert_eq!(puback.first, 0x40); + assert!( + puback.body.len() <= 3, + "[MQTT-3.4.2-2/3] client added PUBACK properties" + ); + match puback.packet() { + Packet::PubAck(ack) => assert!( + mqtt5::packet::is_valid_publish_ack_reason_code(ack.reason_code), + "[MQTT-3.4.2-1] PUBACK reason code {:?}", + ack.reason_code + ), + _ => unreachable!(), + } + + stream + .write_all(&raw_publish(2, false, "t/in", Some(12), &[], b"q2")) + .await + .unwrap(); + let pubrec = expect_kind(&mut stream, PUBREC, MEDIUM).await; + assert_eq!(pubrec.packet_id(), 12, "[MQTT-3.3.4-1] PUBREC id"); + assert!( + pubrec.body.len() <= 3, + "[MQTT-3.5.2-2/3] client added PUBREC properties" + ); + assert!( + mqtt5::packet::is_valid_publish_ack_reason_code( + ReasonCode::from_u8(pubrec.reason_code()).expect("known reason code") + ), + "[MQTT-3.5.2-1] PUBREC reason code 0x{:02X}", + pubrec.reason_code() + ); + stream.write_all(&ack_bytes(0x62, 12)).await.unwrap(); + let pubcomp = expect_kind(&mut stream, PUBCOMP, MEDIUM).await; + assert_eq!(pubcomp.packet_id(), 12, "[MQTT-4.3.3-11] PUBCOMP id"); + assert!( + matches!(pubcomp.reason_code(), 0x00 | 0x92), + "[MQTT-3.7.2-1] PUBCOMP reason code 0x{:02X}", + pubcomp.reason_code() + ); + assert!( + pubcomp.body.len() <= 3, + "[MQTT-3.7.2-2/3] client added PUBCOMP properties" + ); + + assert_eq!(drain(&mut rx, SHORT).await.len(), 2); +} + +#[tokio::test] +async fn mqtt_4_3_3_10_duplicate_qos2_before_pubrel_delivered_once_then_id_reusable() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-dupq2"), &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/q2", qos_options(QoS::ExactlyOnce)).await; + + stream + .write_all(&raw_publish(2, false, "t/q2", Some(5), &[], b"first")) + .await + .unwrap(); + assert_eq!( + expect_kind(&mut stream, PUBREC, MEDIUM).await.packet_id(), + 5 + ); + stream + .write_all(&raw_publish(2, true, "t/q2", Some(5), &[], b"first")) + .await + .unwrap(); + assert_eq!( + expect_kind(&mut stream, PUBREC, MEDIUM).await.packet_id(), + 5, + "[MQTT-4.3.3-10] duplicate before PUBREL must be re-acknowledged with PUBREC" + ); + assert_eq!( + drain(&mut rx, SHORT).await.len(), + 1, + "[MQTT-4.3.3-10] duplicate QoS2 PUBLISH before PUBREL delivered more than once" + ); + + stream.write_all(&ack_bytes(0x62, 5)).await.unwrap(); + assert_eq!( + expect_kind(&mut stream, PUBCOMP, MEDIUM).await.packet_id(), + 5 + ); + + stream + .write_all(&raw_publish(2, false, "t/q2", Some(5), &[], b"second")) + .await + .unwrap(); + assert_eq!( + expect_kind(&mut stream, PUBREC, MEDIUM).await.packet_id(), + 5 + ); + let reused = drain(&mut rx, SHORT).await; + assert_eq!( + reused.len(), + 1, + "[MQTT-4.3.3-12] PUBLISH reusing a completed QoS2 id must be a new message" + ); + assert_eq!(reused[0].payload, b"second"); +} + +#[tokio::test] +async fn mqtt_4_3_2_5_qos1_id_reuse_after_puback_is_new_message() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-q1reuse"), &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/q1", qos_options(QoS::AtLeastOnce)).await; + for (dup, payload) in [(false, &b"a"[..]), (true, &b"b"[..])] { + stream + .write_all(&raw_publish(1, dup, "t/q1", Some(9), &[], payload)) + .await + .unwrap(); + assert_eq!( + expect_kind(&mut stream, PUBACK, MEDIUM).await.packet_id(), + 9 + ); + } + assert_eq!( + drain(&mut rx, SHORT).await.len(), + 2, + "[MQTT-4.3.2-5] QoS1 PUBLISH reusing an acknowledged id must be treated as new" + ); +} + +#[tokio::test] +async fn mqtt_3_7_2_1_pubrel_for_unknown_id_is_completed() { + let (_client, mut stream, _listener) = + connect_client(base_options("conf-b-pubrel"), &plain_connack()).await; + stream.write_all(&ack_bytes(0x62, 77)).await.unwrap(); + let pubcomp = expect_kind(&mut stream, PUBCOMP, MEDIUM).await; + assert_eq!(pubcomp.packet_id(), 77, "[MQTT-4.3.3-11] PUBCOMP id"); + assert!( + matches!(pubcomp.reason_code(), 0x00 | 0x92), + "[MQTT-3.7.2-1] PUBCOMP reason code 0x{:02X}", + pubcomp.reason_code() + ); +} + +#[tokio::test] +async fn mqtt_3_3_1_4_inbound_qos3_publish_is_malformed() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-qos3"), &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/q3", qos_options(QoS::ExactlyOnce)).await; + stream + .write_all(&raw_publish(3, false, "t/q3", Some(3), &[], b"bad")) + .await + .unwrap(); + let delivered = drain(&mut rx, SHORT).await.len(); + let end = termination(&mut stream, MEDIUM).await; + assert_eq!( + delivered, 0, + "[MQTT-3.3.1-4] QoS 3 PUBLISH was delivered to the application" + ); + assert!( + matches!(end, Termination::Closed | Termination::Disconnect(0x80..)), + "[MQTT-3.3.1-4 / §4.13] client did not close the Network Connection after a malformed QoS 3 PUBLISH: {end:?}" + ); +} + +fn alias_client_options(client_id: &str) -> ConnectOptions { + let mut options = base_options(client_id); + options.properties.topic_alias_maximum = Some(2); + options +} + +#[tokio::test] +async fn mqtt_3_3_2_10_client_accepts_topic_alias_within_its_maximum() { + let (listener, url) = bind().await; + let client = MqttClient::with_options(alias_client_options("conf-b-alias-ok")); + let connecting = client.clone(); + let task = tokio::spawn(async move { connecting.connect(&url).await }); + let (mut stream, connect) = accept_session(&listener, &plain_connack()).await; + task.await.unwrap().expect("connect"); + assert_eq!(connect.properties.get_topic_alias_maximum(), Some(2)); + + let mut rx = subscribe(&client, &mut stream, "t/a", qos_options(QoS::AtMostOnce)).await; + stream + .write_all(&raw_publish( + 0, + false, + "t/a", + None, + &topic_alias_prop(1), + b"one", + )) + .await + .unwrap(); + stream + .write_all(&raw_publish( + 0, + false, + "", + None, + &topic_alias_prop(1), + b"two", + )) + .await + .unwrap(); + let received = drain(&mut rx, SHORT).await; + let topics: Vec<&str> = received.iter().map(|m| m.topic.as_str()).collect(); + assert_eq!( + topics, + vec!["t/a", "t/a"], + "[MQTT-3.3.2-10] client did not resolve an inbound Topic Alias (1 <= its maximum 2)" + ); +} + +async fn alias_violation(alias: u16, client_id: &str) -> (usize, Termination) { + let (client, mut stream, _listener) = + connect_client(alias_client_options(client_id), &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/a", qos_options(QoS::AtMostOnce)).await; + stream + .write_all(&raw_publish( + 0, + false, + "t/a", + None, + &topic_alias_prop(alias), + b"x", + )) + .await + .unwrap(); + let delivered = drain(&mut rx, SHORT).await.len(); + (delivered, termination(&mut stream, MEDIUM).await) +} + +#[tokio::test] +async fn mqtt_3_3_2_8_receiver_treats_inbound_topic_alias_zero_as_protocol_error() { + let (delivered, end) = alias_violation(0, "conf-b-alias0").await; + assert!( + delivered == 0 && end == Termination::Disconnect(0x94), + "[MQTT-3.3.2-8 / §3.3.2.3.4] inbound Topic Alias 0 must be a Protocol Error (DISCONNECT 0x94); delivered={delivered} termination={end:?}" + ); +} + +#[tokio::test] +async fn mqtt_3_3_2_11_receiver_treats_topic_alias_above_its_maximum_as_protocol_error() { + let (delivered, end) = alias_violation(3, "conf-b-alias3").await; + assert!( + delivered == 0 && end == Termination::Disconnect(0x94), + "[MQTT-3.3.2-11 / §3.3.2.3.4] inbound Topic Alias 3 > client maximum 2 must be a Protocol Error (DISCONNECT 0x94); delivered={delivered} termination={end:?}" + ); +} + +#[tokio::test] +async fn crosscheck3_inbound_subscription_identifier_zero_is_protocol_error() { + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-inbound-subid0"), &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/sid", qos_options(QoS::AtMostOnce)).await; + stream + .write_all(&raw_publish( + 0, + false, + "t/sid", + None, + &subscription_id_prop(0), + b"x", + )) + .await + .unwrap(); + let delivered = drain(&mut rx, SHORT).await.len(); + let end = termination(&mut stream, MEDIUM).await; + assert!( + delivered == 0 && matches!(end, Termination::Closed | Termination::Disconnect(0x80..)), + "[§3.3.2.3.8] inbound PUBLISH with Subscription Identifier 0 must be a Protocol Error; delivered={delivered} termination={end:?}" + ); +} + +#[tokio::test] +async fn crosscheck2_mqtt_3_3_4_9_server_exceeding_client_receive_maximum_gets_disconnect_0x93() { + let options = base_options("conf-b-rm-in").with_receive_maximum(1); + let (client, mut stream, _listener) = connect_client(options, &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/rm", qos_options(QoS::ExactlyOnce)).await; + stream + .write_all(&raw_publish(2, false, "t/rm", Some(1), &[], b"1")) + .await + .unwrap(); + assert_eq!( + expect_kind(&mut stream, PUBREC, MEDIUM).await.packet_id(), + 1 + ); + stream + .write_all(&raw_publish(2, false, "t/rm", Some(2), &[], b"2")) + .await + .unwrap(); + let end = termination(&mut stream, MEDIUM).await; + let delivered = drain(&mut rx, SHORT).await.len(); + assert!( + end == Termination::Disconnect(0x93), + "[MQTT-3.3.4-9 / §3.3.4] second unacknowledged QoS2 PUBLISH with client Receive Maximum 1 must draw DISCONNECT 0x93; termination={end:?} delivered={delivered}" + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_24_prose_inbound_packet_above_client_maximum_packet_size() { + let mut options = base_options("conf-b-maxpkt"); + options.properties.maximum_packet_size = Some(128); + let (client, mut stream, _listener) = connect_client(options, &plain_connack()).await; + let mut rx = subscribe(&client, &mut stream, "t/big", qos_options(QoS::AtMostOnce)).await; + stream + .write_all(&raw_publish(0, false, "t/big", None, &[], &[b'x'; 1000])) + .await + .unwrap(); + let delivered = drain(&mut rx, SHORT).await.len(); + let end = termination(&mut stream, MEDIUM).await; + assert!( + delivered == 0 && matches!(end, Termination::Closed | Termination::Disconnect(0x95)), + "[§3.1.2.11.4] packet larger than client Maximum Packet Size 128 must be a Protocol Error (DISCONNECT 0x95); delivered={delivered} termination={end:?}" + ); +} + +#[tokio::test] +async fn crosscheck1_mqtt_3_3_4_7_session_resume_with_smaller_receive_maximum() { + let (listener, url) = bind().await; + let client = MqttClient::with_options(resumable_options("conf-b-resume-rm")); + let connecting = client.clone(); + let connect_task = tokio::spawn(async move { connecting.connect(&url).await }); + let (mut first, _) = accept_session( + &listener, + &connack(false, |c| c.properties.set_receive_maximum(2)), + ) + .await; + connect_task.await.unwrap().expect("connect"); + + let mut old = Vec::new(); + for topic in ["t/old/1", "t/old/2"] { + let publisher = client.clone(); + old.push(tokio::spawn(async move { + publisher.publish_qos1(topic, b"old".to_vec()).await + })); + } + let (frames, _) = collect_frames(&mut first, SHORT).await; + let old_ids: Vec = frames + .iter() + .filter(|f| f.kind() == PUBLISH) + .map(Frame::packet_id) + .collect(); + assert_eq!( + old_ids.len(), + 2, + "both QoS1 publishes reach the wire on RM=2" + ); + drop(first); + wait_until_disconnected(&client).await; + + let (mut second, connect) = accept_session( + &listener, + &connack(true, |c| c.properties.set_receive_maximum(1)), + ) + .await; + assert!(!connect.clean_start, "reconnect uses Clean Start 0"); + for _ in 0..100 { + if client.is_connected().await { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + + let mut new = Vec::new(); + for topic in ["t/new/1", "t/new/2"] { + let publisher = client.clone(); + new.push(tokio::spawn(async move { + publisher.publish_qos1(topic, b"new".to_vec()).await + })); + } + + let (frames, closed) = collect_frames(&mut second, Duration::from_millis(1500)).await; + assert!(!closed, "client dropped the resumed connection"); + let unacked: Vec<(u16, bool)> = frames + .iter() + .filter(|f| f.kind() == PUBLISH) + .map(|f| (f.packet_id(), f.dup())) + .collect(); + assert!( + unacked.len() <= 1, + "[MQTT-3.3.4-7 / MQTT-4.9.0-2] client had {} unacknowledged QoS1 PUBLISH on the wire with server Receive Maximum 1: {unacked:?}", + unacked.len() + ); + + let mut acked = 0; + let mut pending = unacked; + let deadline = tokio::time::Instant::now() + Duration::from_secs(5); + while acked < old_ids.len() + 2 && tokio::time::Instant::now() < deadline { + for (id, _) in pending.drain(..) { + second.write_all(&ack_bytes(0x40, id)).await.unwrap(); + acked += 1; + } + let left = deadline.saturating_duration_since(tokio::time::Instant::now()); + if let Next::Frame(frame) = next_frame(&mut second, left).await { + if frame.kind() == PUBLISH { + pending.push((frame.packet_id(), frame.dup())); + } + } + } + for handle in new { + let outcome = tokio::time::timeout(Duration::from_secs(5), handle) + .await + .expect("[MQTT-4.9.0-2] new publish deadlocked after session resume"); + let result = outcome.expect("publish task must not panic"); + assert!( + result.is_ok(), + "new publish after resume failed: {result:?}" + ); + } + for handle in old { + let outcome = tokio::time::timeout(Duration::from_secs(1), handle).await; + if let Ok(joined) = outcome { + assert!( + !joined.is_err_and(|e| e.is_panic()), + "publish task panicked" + ); + } + } +} + +#[tokio::test] +async fn mqtt_4_4_0_1_mqtt_3_3_1_1_unacked_publish_resent_with_dup_on_session_resume() { + let (listener, url) = bind().await; + let client = MqttClient::with_options(resumable_options("conf-b-resend")); + let connecting = client.clone(); + let connect_task = tokio::spawn(async move { connecting.connect(&url).await }); + let (mut first, _) = accept_session(&listener, &plain_connack()).await; + connect_task.await.unwrap().expect("connect"); + + let publisher = client.clone(); + let pending = + tokio::spawn(async move { publisher.publish_qos1("t/resend", b"m".to_vec()).await }); + let original = expect_kind(&mut first, PUBLISH, MEDIUM).await; + drop(first); + wait_until_disconnected(&client).await; + + let (mut second, _) = accept_session(&listener, &connack(true, |_| {})).await; + let resent = match next_frame(&mut second, MEDIUM).await { + Next::Frame(frame) if frame.kind() == PUBLISH => Some(frame), + _ => None, + }; + pending.abort(); + let resent = resent.unwrap_or_else(|| { + panic!( + "[MQTT-4.4.0-1] client did not resend unacknowledged QoS1 PUBLISH id {} after reconnecting with Clean Start 0 and Session Present 1", + original.packet_id() + ) + }); + assert_eq!( + resent.packet_id(), + original.packet_id(), + "[MQTT-4.4.0-1] original packet id" + ); + assert!( + resent.dup(), + "[MQTT-3.3.1-1] re-delivered PUBLISH must have DUP=1" + ); +} + +#[tokio::test] +async fn mqtt_4_9_0_1_send_quota_reinitialized_on_new_connection() { + let (listener, url) = bind().await; + let options = ConnectOptions::new("conf-b-quota") + .with_automatic_reconnect(true) + .with_reconnect_delay(Duration::from_millis(100), Duration::from_millis(200)); + let client = MqttClient::with_options(options); + let connecting = client.clone(); + let connect_task = tokio::spawn(async move { connecting.connect(&url).await }); + let rm1 = connack(false, |c| c.properties.set_receive_maximum(1)); + let (mut first, _) = accept_session(&listener, &rm1).await; + connect_task.await.unwrap().expect("connect"); + + let timed_out = client + .publish_qos1("t/quota", b"never-acked".to_vec()) + .await; + assert!( + matches!(timed_out, Err(MqttError::Timeout)), + "unacknowledged publish times out: {timed_out:?}" + ); + let _ = collect_frames(&mut first, Duration::from_millis(50)).await; + drop(first); + wait_until_disconnected(&client).await; + + let (mut second, _) = accept_session(&listener, &rm1).await; + for _ in 0..100 { + if client.is_connected().await { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + let publisher = client.clone(); + let fresh = + tokio::spawn(async move { publisher.publish_qos1("t/quota", b"fresh".to_vec()).await }); + let sent = matches!( + next_frame(&mut second, Duration::from_secs(2)).await, + Next::Frame(ref f) if f.kind() == PUBLISH + ); + fresh.abort(); + assert!( + sent, + "[MQTT-4.9.0-1] new connection (Session Present 0, Receive Maximum 1) started with a send quota of 0: a permit held by a timed-out publish on the previous connection was never reclaimed" + ); +} + +async fn queued_flush(client_id: &str) -> Vec { + let (listener, url) = bind().await; + let options = ConnectOptions::new(client_id) + .with_clean_start(false) + .with_session_expiry_interval(300) + .with_automatic_reconnect(true) + .with_reconnect_delay(Duration::from_millis(1500), Duration::from_millis(2000)); + let client = MqttClient::with_options(options); + client.set_queue_on_disconnect(true).await; + let connecting = client.clone(); + let connect_task = tokio::spawn(async move { connecting.connect(&url).await }); + let (first, _) = accept_session(&listener, &plain_connack()).await; + connect_task.await.unwrap().expect("connect"); + drop(first); + wait_until_disconnected(&client).await; + + for i in 0..3u8 { + let queued = client + .publish_qos1("t/queued", vec![i]) + .await + .expect("publish while disconnected is queued"); + assert!(queued.packet_id().is_some()); + } + + let (mut second, _) = accept_session( + &listener, + &connack(true, |c| c.properties.set_receive_maximum(1)), + ) + .await; + let (frames, _) = collect_frames(&mut second, Duration::from_millis(1500)).await; + frames.into_iter().filter(|f| f.kind() == PUBLISH).collect() +} + +#[tokio::test] +async fn mqtt_3_3_4_7_queued_messages_flushed_within_server_receive_maximum() { + let publishes = queued_flush("conf-b-queue-rm").await; + assert!( + !publishes.is_empty(), + "queued messages were never sent after reconnect" + ); + assert!( + publishes.len() <= 1, + "[MQTT-3.3.4-7] client flushed {} unacknowledged queued QoS1 PUBLISH packets with server Receive Maximum 1", + publishes.len() + ); +} + +#[tokio::test] +async fn mqtt_4_3_2_2_queued_message_first_transmission_has_dup_zero() { + let publishes = queued_flush("conf-b-queue-dup").await; + let first = publishes + .first() + .expect("queued messages were never sent after reconnect"); + assert!( + !first.dup(), + "[MQTT-4.3.2-2] queued QoS1 message sent for the first time with DUP=1 (packet id {})", + first.packet_id() + ); +} + +#[tokio::test] +async fn mqtt_3_3_4_8_disconnect_not_delayed_by_exhausted_send_quota() { + let rm1 = connack(false, |c| c.properties.set_receive_maximum(1)); + let (client, mut stream, _listener) = + connect_client(base_options("conf-b-nodelay"), &rm1).await; + + let mut blocked = Vec::new(); + for i in 0..2u8 { + let publisher = client.clone(); + blocked.push(tokio::spawn(async move { + publisher.publish_qos1("t/hold", vec![i]).await + })); + } + let _ = expect_kind(&mut stream, PUBLISH, MEDIUM).await; + + let subscriber = client.clone(); + let sub_task = tokio::spawn(async move { subscriber.subscribe("t/other", |_| {}).await }); + let sub = expect_kind(&mut stream, SUBSCRIBE, Duration::from_secs(2)).await; + stream + .write_all(&suback_bytes(sub.packet_id(), 0)) + .await + .unwrap(); + sub_task + .await + .unwrap() + .expect("subscribe while quota is exhausted"); + + let disconnecting = client.clone(); + let disconnect_task = tokio::spawn(async move { disconnecting.disconnect().await }); + let end = termination(&mut stream, Duration::from_secs(2)).await; + disconnect_task.abort(); + for handle in blocked { + handle.abort(); + } + assert!( + matches!(end, Termination::Disconnect(_)), + "[MQTT-3.3.4-8] DISCONNECT was delayed while the send quota was exhausted (nothing within 2s): {end:?}" + ); +} + +fn deferred_options(client_id: &str, receive_maximum: u16) -> ConnectOptions { + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_session_expiry_interval(300) + .with_receive_maximum(receive_maximum) + .with_deferred_ack(true) + .with_automatic_reconnect(false) +} + +async fn deferred_reject_code(qos: QoS, reject_with: ReasonCode) -> Frame { + let (client, mut stream, _listener) = + connect_client(deferred_options("conf-b-reject", 10), &plain_connack()).await; + let subscriber = client.clone(); + let task = tokio::spawn(async move { + subscriber + .subscribe_with_ack("t/rej", qos_options(qos), move |_, token: AckToken| { + token.reject(reject_with); + }) + .await + }); + let sub = expect_kind(&mut stream, SUBSCRIBE, MEDIUM).await; + stream + .write_all(&suback_bytes(sub.packet_id(), qos as u8)) + .await + .unwrap(); + task.await.unwrap().expect("subscribe_with_ack"); + stream + .write_all(&raw_publish(qos as u8, false, "t/rej", Some(21), &[], b"x")) + .await + .unwrap(); + let kind = if qos == QoS::AtLeastOnce { + PUBACK + } else { + PUBREC + }; + expect_kind(&mut stream, kind, MEDIUM).await +} + +#[tokio::test] +async fn mqtt_3_4_2_1_deferred_reject_puback_reason_code_must_be_valid() { + let ack = deferred_reject_code(QoS::AtLeastOnce, ReasonCode::ServerBusy).await; + let code = ReasonCode::from_u8(ack.reason_code()).expect("known reason code"); + assert!( + mqtt5::packet::is_valid_publish_ack_reason_code(code), + "[MQTT-3.4.2-1] AckToken::reject put PUBACK reason code {code:?} (0x{:02X}) on the wire, which is not a PUBACK Reason Code", + ack.reason_code() + ); +} + +#[tokio::test] +async fn mqtt_3_5_2_1_deferred_reject_pubrec_reason_code_must_be_valid() { + let ack = deferred_reject_code(QoS::ExactlyOnce, ReasonCode::ServerBusy).await; + let code = ReasonCode::from_u8(ack.reason_code()).expect("known reason code"); + assert!( + mqtt5::packet::is_valid_publish_ack_reason_code(code), + "[MQTT-3.5.2-1] AckToken::reject put PUBREC reason code {code:?} (0x{:02X}) on the wire, which is not a PUBREC Reason Code", + ack.reason_code() + ); +} + +#[tokio::test] +async fn mqtt_4_4_0_1_deferred_ack_redelivery_at_receive_maximum_after_resume() { + let (listener, url) = bind().await; + let mut options = deferred_options("conf-b-deferred-resume", 1); + options.reconnect_config.enabled = true; + options.reconnect_config.initial_delay = Duration::from_millis(100); + options.reconnect_config.max_delay = Duration::from_millis(200); + let client = MqttClient::with_options(options); + let connecting = client.clone(); + let connect_task = tokio::spawn(async move { connecting.connect(&url).await }); + let (mut first, _) = accept_session(&listener, &plain_connack()).await; + connect_task.await.unwrap().expect("connect"); + + let held: Arc>> = Arc::new(Mutex::new(Vec::new())); + let deliveries = Arc::new(Mutex::new(Vec::>::new())); + let subscriber = client.clone(); + let (held_cb, deliveries_cb) = (Arc::clone(&held), Arc::clone(&deliveries)); + let sub_task = tokio::spawn(async move { + subscriber + .subscribe_with_ack( + "t/deferred", + qos_options(QoS::AtLeastOnce), + move |publish, token| { + deliveries_cb.lock().unwrap().push(publish.payload.to_vec()); + held_cb.lock().unwrap().push(token); + }, + ) + .await + }); + let sub = expect_kind(&mut first, SUBSCRIBE, MEDIUM).await; + first + .write_all(&suback_bytes(sub.packet_id(), 1)) + .await + .unwrap(); + sub_task.await.unwrap().expect("subscribe_with_ack"); + + first + .write_all(&raw_publish(1, false, "t/deferred", Some(1), &[], b"held")) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(300)).await; + assert_eq!( + held.lock().unwrap().len(), + 1, + "first delivery holds its token" + ); + drop(first); + wait_until_disconnected(&client).await; + + let (mut second, _) = accept_session(&listener, &connack(true, |_| {})).await; + tokio::time::sleep(Duration::from_millis(200)).await; + second + .write_all(&raw_publish(1, true, "t/deferred", Some(1), &[], b"held")) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(300)).await; + let tokens: Vec = held.lock().unwrap().drain(..).collect(); + for token in tokens { + token.ack(); + } + let _ = collect_frames(&mut second, Duration::from_millis(300)).await; + + second + .write_all(&raw_publish(1, false, "t/deferred", Some(2), &[], b"next")) + .await + .unwrap(); + tokio::time::sleep(Duration::from_millis(500)).await; + let got_next = deliveries + .lock() + .unwrap() + .iter() + .any(|payload| payload == b"next"); + assert!( + got_next, + "[MQTT-4.4.0-1 / §4.9] a resumed-session DUP redelivery of a still-held message at client Receive Maximum 1 broke the connection: later PUBLISH id 2 was never delivered" + ); +} + +#[tokio::test] +async fn crosscheck4_topic_alias_boundary_no_panic_via_public_api() { + let outcome = tokio::spawn(async move { + let mut inbound = TopicAliasManager::new(u16::MAX); + inbound + .register_alias(u16::MAX, "t/max") + .expect("alias == maximum is valid"); + let mut outbound = TopicAliasManager::new(u16::MAX); + let mut last = None; + for i in 0..u32::from(u16::MAX) { + last = outbound.get_or_create_alias(&format!("t/{i}")); + } + last + }) + .await; + match outcome { + Ok(last) => assert_eq!(last, Some(u16::MAX)), + Err(e) => panic!( + "[cross-check 4] mqtt5::session::TopicAliasManager panicked assigning outbound Topic Alias == Topic Alias Maximum (65535): {e}" + ), + } + + let (client, mut stream, _listener) = connect_client( + base_options("conf-b-alias-boundary-wire"), + &connack(false, |c| c.properties.set_topic_alias_maximum(u16::MAX)), + ) + .await; + let properties = PublishProperties { + topic_alias: Some(u16::MAX), + ..Default::default() + }; + client + .publish_with_options("t/max", b"x".to_vec(), qos0_with(properties)) + .await + .expect("publish with alias == server maximum"); + let frame = expect_kind(&mut stream, PUBLISH, MEDIUM).await; + assert_eq!(frame.publish().topic_alias(), Some(u16::MAX)); +} + +#[tokio::test] +async fn mqtt_3_2_2_5_queued_acks_discarded_when_session_not_present() { + let (listener, url) = bind().await; + let mut options = deferred_options("conf-b-sp0-acks", 4); + options.reconnect_config.enabled = true; + options.reconnect_config.initial_delay = Duration::from_millis(100); + options.reconnect_config.max_delay = Duration::from_millis(200); + let client = MqttClient::with_options(options); + let connecting = client.clone(); + let connect_task = tokio::spawn(async move { connecting.connect(&url).await }); + let (mut first, _) = accept_session(&listener, &plain_connack()).await; + connect_task.await.unwrap().expect("connect"); + + let held: Arc>> = Arc::new(Mutex::new(Vec::new())); + let subscriber = client.clone(); + let held_cb = Arc::clone(&held); + let sub_task = tokio::spawn(async move { + subscriber + .subscribe_with_ack( + "t/deferred", + qos_options(QoS::AtLeastOnce), + move |_, token| held_cb.lock().unwrap().push(token), + ) + .await + }); + let sub = expect_kind(&mut first, SUBSCRIBE, MEDIUM).await; + first + .write_all(&suback_bytes(sub.packet_id(), 1)) + .await + .unwrap(); + sub_task.await.unwrap().expect("subscribe_with_ack"); + let _auto = subscribe(&client, &mut first, "t/auto", qos_options(QoS::AtLeastOnce)).await; + + first + .write_all(&raw_publish(1, false, "t/deferred", Some(1), &[], b"held")) + .await + .unwrap(); + first + .write_all(&raw_publish(1, false, "t/auto", Some(2), &[], b"auto")) + .await + .unwrap(); + let (before, _) = collect_frames(&mut first, Duration::from_millis(400)).await; + assert!( + before.iter().all(|f| f.kind() != PUBACK), + "setup: the automatic PUBACK for id 2 is held behind the unresolved token" + ); + assert_eq!(held.lock().unwrap().len(), 1, "setup: token held"); + drop(first); + wait_until_disconnected(&client).await; + + let (mut second, _) = accept_session(&listener, &plain_connack()).await; + for _ in 0..100 { + if client.is_connected().await { + break; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + let tokens: Vec = held.lock().unwrap().drain(..).collect(); + for token in tokens { + token.ack(); + } + let (frames, _) = collect_frames(&mut second, Duration::from_millis(800)).await; + let pubacks: Vec = frames + .iter() + .filter(|f| f.kind() == PUBACK) + .map(Frame::packet_id) + .collect(); + assert!( + pubacks.is_empty(), + "[MQTT-3.2.2-5] acknowledgements from the discarded session were sent on a Session Present=0 connection: {pubacks:?}" + ); +} diff --git a/crates/mqtt5/tests/conf_client_c.rs b/crates/mqtt5/tests/conf_client_c.rs new file mode 100644 index 00000000..1ec7774b --- /dev/null +++ b/crates/mqtt5/tests/conf_client_c.rs @@ -0,0 +1,961 @@ +use mqtt5::time::Duration; +use mqtt5::{ + AckToken, AuthHandler, AuthResponse, ConnectOptions, MqttClient, QoS, SubscribeOptions, +}; +use std::future::Future; +use std::pin::Pin; +use std::sync::atomic::{AtomicU32, Ordering}; +use std::sync::{Arc, Mutex}; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; + +const TEST_TIMEOUT: Duration = Duration::from_secs(40); +const DISCONNECT_REASON_TABLE: [u8; 29] = [ + 0x00, 0x04, 0x80, 0x81, 0x82, 0x83, 0x87, 0x89, 0x8B, 0x8D, 0x8E, 0x8F, 0x90, 0x93, 0x94, 0x95, + 0x96, 0x97, 0x98, 0x99, 0x9A, 0x9B, 0x9C, 0x9D, 0x9E, 0x9F, 0xA0, 0xA1, 0xA2, +]; +const AUTH_REASON_TABLE: [u8; 3] = [0x00, 0x18, 0x19]; + +#[derive(Debug, Clone, PartialEq, Eq)] +struct Raw { + header: u8, + body: Vec, +} + +impl Raw { + fn kind(&self) -> u8 { + self.header >> 4 + } + + fn ack_id(&self) -> u16 { + u16::from_be_bytes([self.body[0], self.body[1]]) + } + + fn reason(&self) -> u8 { + match self.kind() { + 14 | 15 => self.body.first().copied().unwrap_or(0), + _ => self.body.get(2).copied().unwrap_or(0), + } + } + + fn publish_qos(&self) -> u8 { + (self.header >> 1) & 0x03 + } + + fn publish_dup(&self) -> bool { + self.header & 0x08 != 0 + } + + fn publish_id(&self) -> Option { + let topic_len = usize::from(u16::from_be_bytes([self.body[0], self.body[1]])); + (self.publish_qos() > 0) + .then(|| u16::from_be_bytes([self.body[2 + topic_len], self.body[3 + topic_len]])) + } + + fn describe(&self) -> String { + match self.kind() { + 3 => format!( + "PUBLISH(qos={},dup={},id={:?})", + self.publish_qos(), + self.publish_dup(), + self.publish_id() + ), + 4 => format!("PUBACK({})", self.ack_id()), + 5 => format!("PUBREC({})", self.ack_id()), + 6 => format!("PUBREL({})", self.ack_id()), + 7 => format!("PUBCOMP({})", self.ack_id()), + 8 => "SUBSCRIBE".to_string(), + 12 => "PINGREQ".to_string(), + 14 => format!("DISCONNECT(0x{:02X})", self.reason()), + 15 => format!("AUTH(0x{:02X})", self.reason()), + other => format!("type{other}"), + } + } +} + +enum Next { + Packet(Raw), + Eof, + Timeout, +} + +fn encode(header: u8, body: &[u8]) -> Vec { + let mut out = vec![header]; + let mut len = body.len(); + loop { + let mut byte = u8::try_from(len % 128).expect("fits"); + len /= 128; + if len > 0 { + byte |= 0x80; + } + out.push(byte); + if len == 0 { + break; + } + } + out.extend_from_slice(body); + out +} + +fn connack(session_present: bool) -> Vec { + encode(0x20, &[u8::from(session_present), 0x00, 0x00]) +} + +fn server_publish(qos: u8, dup: bool, packet_id: u16, topic: &str, payload: &[u8]) -> Vec { + let topic_len = u16::try_from(topic.len()).expect("topic fits"); + let mut body = topic_len.to_be_bytes().to_vec(); + body.extend_from_slice(topic.as_bytes()); + if qos > 0 { + body.extend_from_slice(&packet_id.to_be_bytes()); + } + body.push(0x00); + body.extend_from_slice(payload); + encode(0x30 | (u8::from(dup) << 3) | (qos << 1), &body) +} + +fn ack(header: u8, packet_id: u16) -> Vec { + encode(header, &packet_id.to_be_bytes()) +} + +fn server_auth(reason: u8, method: &str) -> Vec { + let method_len = u16::try_from(method.len()).expect("method fits"); + let mut props = vec![0x15]; + props.extend_from_slice(&method_len.to_be_bytes()); + props.extend_from_slice(method.as_bytes()); + let mut body = vec![reason, u8::try_from(props.len()).expect("props fit")]; + body.extend_from_slice(&props); + encode(0xF0, &body) +} + +struct Wire { + stream: TcpStream, + clean_start: bool, +} + +impl Wire { + async fn raw_read(&mut self) -> Option { + let mut first = [0u8; 1]; + if self.stream.read_exact(&mut first).await.is_err() { + return None; + } + let mut len = 0usize; + let mut shift = 0; + loop { + let mut b = [0u8; 1]; + if self.stream.read_exact(&mut b).await.is_err() { + return None; + } + len |= usize::from(b[0] & 0x7F) << shift; + shift += 7; + if b[0] & 0x80 == 0 { + break; + } + } + let mut body = vec![0u8; len]; + if self.stream.read_exact(&mut body).await.is_err() { + return None; + } + Some(Raw { + header: first[0], + body, + }) + } + + async fn next(&mut self, wait: Duration) -> Next { + loop { + match tokio::time::timeout(wait, self.raw_read()).await { + Err(_) => return Next::Timeout, + Ok(None) => return Next::Eof, + Ok(Some(raw)) if raw.kind() == 12 => { + self.send(&[0xD0, 0x00]).await; + } + Ok(Some(raw)) if raw.kind() == 8 => { + let mut suback = raw.body[0..2].to_vec(); + suback.extend_from_slice(&[0x00, 0x02]); + self.send(&encode(0x90, &suback)).await; + } + Ok(Some(raw)) => return Next::Packet(raw), + } + } + } + + async fn expect(&mut self, what: &str) -> Raw { + match self.next(Duration::from_secs(5)).await { + Next::Packet(raw) => raw, + Next::Eof => panic!("connection closed while waiting for {what}"), + Next::Timeout => panic!("timed out waiting for {what}"), + } + } + + async fn collect(&mut self, wait: Duration) -> (Vec, bool) { + let mut out = Vec::new(); + loop { + match self.next(wait).await { + Next::Packet(raw) => out.push(raw), + Next::Eof => return (out, true), + Next::Timeout => return (out, false), + } + } + } + + async fn send(&mut self, bytes: &[u8]) { + self.stream + .write_all(bytes) + .await + .expect("fake broker write"); + self.stream.flush().await.expect("fake broker flush"); + } +} + +async fn accept(listener: &TcpListener, session_present: bool) -> Wire { + let (stream, _) = tokio::time::timeout(Duration::from_secs(10), listener.accept()) + .await + .expect("client did not connect in time") + .expect("accept"); + let mut wire = Wire { + stream, + clean_start: false, + }; + let connect = wire.raw_read().await.expect("CONNECT"); + assert_eq!(connect.kind(), 1, "first packet must be CONNECT"); + wire.clean_start = connect.body[7] & 0x02 != 0; + wire.send(&connack(session_present)).await; + wire +} + +fn url(listener: &TcpListener) -> String { + format!("mqtt://{}", listener.local_addr().expect("addr")) +} + +fn base_options(name: &str) -> ConnectOptions { + ConnectOptions::new(name) + .with_keep_alive(Duration::from_secs(60)) + .with_automatic_reconnect(false) +} + +fn persistent_options(name: &str) -> ConnectOptions { + ConnectOptions::new(name) + .with_keep_alive(Duration::from_secs(60)) + .with_clean_start(false) + .with_session_expiry_interval(3600) + .with_reconnect_delay(Duration::from_millis(100), Duration::from_millis(500)) +} + +fn deferred_options(name: &str) -> ConnectOptions { + base_options(name) + .with_clean_start(false) + .with_session_expiry_interval(3600) + .with_receive_maximum(16) + .with_deferred_ack(true) +} + +fn connected( + listener: &TcpListener, + options: ConnectOptions, +) -> Pin + '_>> { + Box::pin(async move { + let client = MqttClient::with_options(options.clone()); + let address = url(listener); + let (result, wire) = tokio::join!( + client.connect_with_options(&address, options), + accept(listener, false) + ); + result.expect("client connect"); + (client, wire) + }) +} + +fn qos_options(qos: QoS) -> SubscribeOptions { + SubscribeOptions { + qos, + ..Default::default() + } +} + +async fn wait_until(cond: impl Fn() -> bool) -> bool { + for _ in 0..200 { + if cond() { + return true; + } + tokio::time::sleep(Duration::from_millis(25)).await; + } + cond() +} + +fn descriptions(packets: &[Raw]) -> Vec { + packets.iter().map(Raw::describe).collect() +} + +async fn with_timeout(fut: Pin>>) { + tokio::time::timeout(TEST_TIMEOUT, fut) + .await + .expect("test exceeded its overall timeout"); +} + +#[tokio::test] +async fn mqtt_3_14_2_1_and_3_14_4_2_user_disconnect_uses_table_code_and_closes() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, base_options("c-disc")).await; + client.disconnect().await.expect("disconnect"); + let (packets, eof) = wire.collect(Duration::from_secs(3)).await; + let disconnect = packets + .iter() + .find(|p| p.kind() == 14) + .expect("MQTT-3.14.2-1: client disconnect() must put a DISCONNECT on the wire"); + assert!( + DISCONNECT_REASON_TABLE.contains(&disconnect.reason()), + "MQTT-3.14.2-1: DISCONNECT reason 0x{:02X} outside table", + disconnect.reason() + ); + assert!( + eof, + "MQTT-3.14.4-2: network connection not closed after DISCONNECT" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_3_14_4_1_no_packet_after_disconnect_when_deferred_ack_resolves_concurrently() { + with_timeout(Box::pin(async { + for round in 0..10 { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = + connected(&listener, deferred_options(&format!("c-race-{round}"))).await; + let tokens: Arc>> = Arc::new(Mutex::new(Vec::new())); + let sink = Arc::clone(&tokens); + let (sub, ()) = tokio::join!( + client.subscribe_with_ack("a", qos_options(QoS::AtLeastOnce), move |_p, t| { + sink.lock().unwrap().push(t); + }), + async { + let _ = wire.next(Duration::from_millis(300)).await; + } + ); + sub.expect("subscribe_with_ack"); + wire.send(&server_publish(1, false, 11, "a", b"x")).await; + assert!(wait_until(|| tokens.lock().unwrap().len() == 1).await); + let token = tokens.lock().unwrap().pop().unwrap(); + token.ack(); + client.disconnect().await.expect("disconnect"); + let (packets, _) = wire.collect(Duration::from_secs(2)).await; + let seen = descriptions(&packets); + if let Some(pos) = packets.iter().position(|p| p.kind() == 14) { + assert!( + pos == packets.len() - 1, + "MQTT-3.14.4-1 violated (round {round}): packets after DISCONNECT: {seen:?}" + ); + } else { + panic!("no DISCONNECT seen (round {round}): {seen:?}"); + } + } + })) + .await; +} + +#[tokio::test] +async fn mqtt_3_14_1_1_disconnect_reserved_bits_client_sends_0x81_then_closes() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (_client, mut wire) = connected(&listener, base_options("c-rsv-disc")).await; + wire.send(&[0xE1, 0x00]).await; + let (packets, eof) = wire.collect(Duration::from_secs(5)).await; + let seen = descriptions(&packets); + assert!( + packets.iter().any(|p| p.kind() == 14 && p.reason() == 0x81), + "MQTT-3.14.1-1 violated: no DISCONNECT 0x81 after DISCONNECT with reserved flags 0x1; saw {seen:?}, socket closed={eof}" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_3_15_1_1_auth_reserved_bits_treated_malformed_and_connection_closed() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, base_options("c-rsv-auth")).await; + wire.send(&[0xF1, 0x00]).await; + let (packets, eof) = wire.collect(Duration::from_secs(5)).await; + let still_connected = client.is_connected().await; + assert!( + eof, + "MQTT-3.15.1-1 violated: AUTH with reserved flags 0x1 did not close the network connection within 5s (client is_connected={still_connected}, saw {:?})", + descriptions(&packets) + ); + })) + .await; +} + +#[tokio::test] +async fn crosscheck4_protocol_error_disconnect_flushed_before_close_with_auto_reconnect() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (_client, mut wire) = connected(&listener, persistent_options("c-proto-err")).await; + wire.send(&[0x36, 0x00]).await; + let (packets, eof) = wire.collect(Duration::from_secs(6)).await; + let seen = descriptions(&packets); + assert!( + eof, + "protocol error: old socket never closed within 6s; saw {seen:?}" + ); + assert!( + packets.iter().any(|p| p.kind() == 14 && p.reason() >= 0x80), + "MQTT-4.13 / cross-check 4: socket closed after malformed PUBLISH (QoS 3) without any DISCONNECT on the wire; saw {seen:?}" + ); + })) + .await; +} + +struct ScriptedAuth { + calls: AtomicU32, +} + +impl AuthHandler for ScriptedAuth { + fn handle_challenge<'a>( + &'a self, + _auth_method: &'a str, + _challenge_data: Option<&'a [u8]>, + ) -> Pin> + Send + 'a>> { + Box::pin(async move { + match self.calls.fetch_add(1, Ordering::SeqCst) { + 0 => Ok(AuthResponse::Continue(b"resp".to_vec())), + 1 => Ok(AuthResponse::Abort("handler refuses".to_string())), + _ => Err(mqtt5::MqttError::AuthenticationFailed), + } + }) + } + + fn initial_response<'a>( + &'a self, + _auth_method: &'a str, + ) -> Pin>>> + Send + 'a>> { + Box::pin(async move { Ok(Some(b"init".to_vec())) }) + } +} + +#[tokio::test] +async fn mqtt_3_14_2_1_and_3_15_2_1_auth_paths_only_emit_table_reason_codes() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let options = base_options("c-auth").with_authentication_method("SCRIPT"); + let client = MqttClient::with_options(options.clone()); + client + .set_auth_handler(ScriptedAuth { + calls: AtomicU32::new(0), + }) + .await; + let address = url(&listener); + let (result, mut wire) = tokio::join!( + client.connect_with_options(&address, options), + accept(&listener, false) + ); + result.expect("connect"); + + client.reauthenticate().await.expect("reauthenticate"); + let reauth = wire.expect("re-auth AUTH").await; + let mut all = vec![reauth]; + wire.send(&server_auth(0x18, "SCRIPT")).await; + all.push(wire.expect("continue AUTH").await); + wire.send(&server_auth(0x18, "SCRIPT")).await; + let (rest, _) = wire.collect(Duration::from_secs(3)).await; + all.extend(rest); + let seen = descriptions(&all); + for p in &all { + match p.kind() { + 14 => assert!( + DISCONNECT_REASON_TABLE.contains(&p.reason()), + "MQTT-3.14.2-1 violated: {seen:?}" + ), + 15 => assert!( + AUTH_REASON_TABLE.contains(&p.reason()), + "MQTT-3.15.2-1 violated: {seen:?}" + ), + _ => {} + } + } + assert_eq!(all[0].kind(), 15, "reauthenticate must send AUTH: {seen:?}"); + assert_eq!(all[0].reason(), 0x19, "re-auth AUTH reason: {seen:?}"); + assert_eq!(all[1].reason(), 0x18, "continue AUTH reason: {seen:?}"); + })) + .await; +} + +async fn collect_acks(wire: &mut Wire, kind: u8, count: usize) -> Vec { + let mut ids = Vec::new(); + while ids.len() < count { + let raw = wire.expect("ack").await; + if raw.kind() == kind { + ids.push(raw.ack_id()); + } + } + ids +} + +#[tokio::test] +async fn mqtt_4_6_0_2_and_4_6_0_3_automatic_acks_follow_publish_order_with_slow_and_panicking_callbacks( +) { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, base_options("c-auto-order")).await; + let (sub, ()) = tokio::join!( + client.subscribe_with_options("t/#", qos_options(QoS::ExactlyOnce), |m| { + assert!(m.payload != b"panic", "callback bug"); + std::thread::sleep(std::time::Duration::from_millis(30)); + }), + async { + let _ = wire.next(Duration::from_millis(300)).await; + } + ); + sub.expect("subscribe"); + let mut burst = Vec::new(); + for id in 1..=6u16 { + let payload: &[u8] = if id == 2 { b"panic" } else { b"x" }; + burst.extend(server_publish(1, false, id, "t/a", payload)); + } + wire.send(&burst).await; + assert_eq!( + collect_acks(&mut wire, 4, 6).await, + vec![1, 2, 3, 4, 5, 6], + "MQTT-4.6.0-2 violated on automatic path" + ); + let mut burst = Vec::new(); + for id in 20..=25u16 { + let payload: &[u8] = if id == 21 { b"panic" } else { b"x" }; + burst.extend(server_publish(2, false, id, "t/b", payload)); + } + wire.send(&burst).await; + assert_eq!( + collect_acks(&mut wire, 5, 6).await, + vec![20, 21, 22, 23, 24, 25], + "MQTT-4.6.0-3 violated on automatic path" + ); + })) + .await; +} + +async fn deferred_reverse_order(qos: u8, ack_kind: u8) -> Vec { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, deferred_options("c-def-rev")).await; + let tokens: Arc>> = Arc::new(Mutex::new(Vec::new())); + let sink = Arc::clone(&tokens); + let sub_qos = if qos == 2 { + QoS::ExactlyOnce + } else { + QoS::AtLeastOnce + }; + let (sub, ()) = tokio::join!( + client.subscribe_with_ack("a", qos_options(sub_qos), move |_p, t| { + sink.lock().unwrap().push(t); + }), + async { + let _ = wire.next(Duration::from_millis(300)).await; + } + ); + sub.expect("subscribe_with_ack"); + let mut burst = server_publish(qos, false, 1, "a", b"first"); + burst.extend(server_publish(qos, false, 2, "a", b"second")); + wire.send(&burst).await; + assert!(wait_until(|| tokens.lock().unwrap().len() == 2).await); + let second = tokens.lock().unwrap().pop().unwrap(); + let first = tokens.lock().unwrap().pop().unwrap(); + second.ack(); + tokio::time::sleep(Duration::from_millis(100)).await; + first.ack(); + collect_acks(&mut wire, ack_kind, 2).await +} + +#[tokio::test] +async fn mqtt_4_6_0_2_deferred_tokens_acked_in_reverse_order() { + with_timeout(Box::pin(async { + let order = Box::pin(deferred_reverse_order(1, 4)).await; + assert_eq!( + order, + vec![1, 2], + "MQTT-4.6.0-2 violated: AckToken API lets PUBACKs reach the wire in resolution order, not PUBLISH arrival order" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_6_0_3_deferred_tokens_acked_in_reverse_order() { + with_timeout(Box::pin(async { + let order = Box::pin(deferred_reverse_order(2, 5)).await; + assert_eq!( + order, + vec![1, 2], + "MQTT-4.6.0-3 violated: AckToken API lets PUBRECs reach the wire in resolution order, not PUBLISH arrival order" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_6_0_2_deferred_acked_immediately_mixed_with_automatic() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, deferred_options("c-mixed")).await; + let (sub, ()) = tokio::join!( + client.subscribe_with_ack("a", qos_options(QoS::AtLeastOnce), |_p, t| t.ack()), + async { + let _ = wire.next(Duration::from_millis(300)).await; + } + ); + sub.expect("subscribe_with_ack"); + let (sub, ()) = tokio::join!( + client.subscribe_with_options("b", qos_options(QoS::AtLeastOnce), |_m| {}), + async { + let _ = wire.next(Duration::from_millis(300)).await; + } + ); + sub.expect("subscribe"); + let mut burst = server_publish(1, false, 1, "a", b"deferred"); + burst.extend(server_publish(1, false, 2, "b", b"automatic")); + wire.send(&burst).await; + assert_eq!( + collect_acks(&mut wire, 4, 2).await, + vec![1, 2], + "MQTT-4.6.0-2 violated: deferred message acked immediately inside its callback still reaches the wire after a later automatic PUBACK" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_3_1_1_4_3_2_2_4_3_3_2_first_send_dup_zero() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, base_options("c-dup0")).await; + let pub_client = client.clone(); + let publisher = tokio::spawn(async move { + pub_client + .publish_qos("t", b"0".to_vec(), QoS::AtMostOnce) + .await?; + pub_client + .publish_qos("t", b"1".to_vec(), QoS::AtLeastOnce) + .await?; + pub_client + .publish_qos("t", b"2".to_vec(), QoS::ExactlyOnce) + .await + }); + let q0 = wire.expect("QoS0 PUBLISH").await; + assert!( + q0.kind() == 3 && q0.publish_qos() == 0 && !q0.publish_dup(), + "MQTT-4.3.1-1: {}", + q0.describe() + ); + let q1 = wire.expect("QoS1 PUBLISH").await; + assert!( + q1.kind() == 3 && q1.publish_qos() == 1 && !q1.publish_dup(), + "MQTT-4.3.2-2: {}", + q1.describe() + ); + wire.send(&ack(0x40, q1.publish_id().unwrap())).await; + let q2 = wire.expect("QoS2 PUBLISH").await; + assert!( + q2.kind() == 3 && q2.publish_qos() == 2 && !q2.publish_dup(), + "MQTT-4.3.3-2: {}", + q2.describe() + ); + let id = q2.publish_id().unwrap(); + wire.send(&ack(0x50, id)).await; + let rel = wire.expect("PUBREL").await; + assert_eq!( + rel.kind(), + 6, + "MQTT-4.3.3-4: expected PUBREL, got {}", + rel.describe() + ); + assert_eq!(rel.header, 0x62, "PUBREL fixed header flags"); + assert_eq!(rel.ack_id(), id, "MQTT-4.3.3-4: PUBREL id mismatch"); + wire.send(&ack(0x70, id)).await; + publisher.await.unwrap().expect("publishes complete"); + let (extra, _) = wire.collect(Duration::from_millis(500)).await; + assert!( + extra.iter().all(|p| p.kind() != 3), + "MQTT-4.3.3-6: PUBLISH re-sent after PUBREL: {:?}", + descriptions(&extra) + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_4_0_2_error_pubrec_gets_no_pubrel() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, base_options("c-pubrec-err")).await; + let pub_client = client.clone(); + let publisher = tokio::spawn(async move { + pub_client + .publish_qos("t", b"2".to_vec(), QoS::ExactlyOnce) + .await + }); + let p = wire.expect("QoS2 PUBLISH").await; + let id = p.publish_id().unwrap(); + let mut rec = id.to_be_bytes().to_vec(); + rec.extend_from_slice(&[0x80, 0x00]); + wire.send(&encode(0x50, &rec)).await; + let (after, _) = wire.collect(Duration::from_secs(1)).await; + assert!( + after.is_empty(), + "MQTT-4.4.0-2 / 4.3.3-4: client answered an error PUBREC: {:?}", + descriptions(&after) + ); + assert!( + publisher.await.unwrap().is_err(), + "error PUBREC must fail the publish" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_3_2_4_4_3_3_8_4_3_3_10_4_3_3_11_receiver_acks_and_dedup() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, base_options("c-recv")).await; + let hits = Arc::new(AtomicU32::new(0)); + let counter = Arc::clone(&hits); + let (sub, ()) = tokio::join!( + client.subscribe_with_options("r", qos_options(QoS::ExactlyOnce), move |_m| { + counter.fetch_add(1, Ordering::SeqCst); + }), + async { + let _ = wire.next(Duration::from_millis(300)).await; + } + ); + sub.expect("subscribe"); + wire.send(&server_publish(1, false, 300, "r", b"q1")).await; + let puback = wire.expect("PUBACK").await; + assert_eq!((puback.kind(), puback.ack_id()), (4, 300), "MQTT-4.3.2-4"); + wire.send(&server_publish(2, false, 301, "r", b"q2")).await; + let pubrec = wire.expect("PUBREC").await; + assert_eq!((pubrec.kind(), pubrec.ack_id()), (5, 301), "MQTT-4.3.3-8"); + wire.send(&server_publish(2, true, 301, "r", b"q2")).await; + let pubrec = wire.expect("PUBREC for duplicate").await; + assert_eq!( + (pubrec.kind(), pubrec.ack_id()), + (5, 301), + "MQTT-4.3.3-10: duplicate must be PUBREC'd" + ); + wire.send(&ack(0x62, 301)).await; + let pubcomp = wire.expect("PUBCOMP").await; + assert_eq!( + (pubcomp.kind(), pubcomp.ack_id()), + (7, 301), + "MQTT-4.3.3-11" + ); + tokio::time::sleep(Duration::from_millis(300)).await; + assert_eq!( + hits.load(Ordering::SeqCst), + 2, + "MQTT-4.3.3-10: QoS2 duplicate delivered to the application" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_5_0_2_publish_without_matching_callback_is_acked() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (_client, mut wire) = connected(&listener, base_options("c-nocb")).await; + wire.send(&server_publish(1, false, 5, "unsubscribed", b"x")) + .await; + let a = wire.expect("PUBACK").await; + assert_eq!((a.kind(), a.ack_id()), (4, 5), "MQTT-4.5.0-2"); + wire.send(&server_publish(2, false, 6, "unsubscribed", b"x")) + .await; + let r = wire.expect("PUBREC").await; + assert_eq!((r.kind(), r.ack_id()), (5, 6), "MQTT-4.5.0-2"); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_4_0_1_and_4_6_0_1_publish_resent_in_order_with_dup_after_session_resume() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, persistent_options("c-resend")).await; + let mut sent = Vec::new(); + for (i, qos) in [QoS::AtLeastOnce, QoS::ExactlyOnce, QoS::AtLeastOnce] + .into_iter() + .enumerate() + { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos("t", vec![u8::try_from(i).unwrap()], qos).await }); + let p = wire.expect("PUBLISH").await; + sent.push(p.publish_id().unwrap()); + } + drop(wire); + let mut wire = accept(&listener, true).await; + assert!(!wire.clean_start, "reconnect must use Clean Start 0"); + let (packets, _) = wire.collect(Duration::from_secs(3)).await; + let resent: Vec<(u16, bool)> = packets + .iter() + .filter(|p| p.kind() == 3) + .map(|p| (p.publish_id().unwrap_or(0), p.publish_dup())) + .collect(); + let expected: Vec<(u16, bool)> = sent.iter().map(|id| (*id, true)).collect(); + assert_eq!( + resent, expected, + "MQTT-4.4.0-1 / 4.6.0-1: unacknowledged PUBLISH {sent:?} must be re-sent in original order with DUP=1 after Session Present=1; wire: {:?}", + descriptions(&packets) + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_6_0_4_pubrel_order_and_resend_after_session_resume() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, persistent_options("c-pubrel")).await; + let mut ids = Vec::new(); + for i in 0..3u8 { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos("t", vec![i], QoS::ExactlyOnce).await }); + ids.push(wire.expect("PUBLISH").await.publish_id().unwrap()); + } + let mut burst = Vec::new(); + for id in &ids { + burst.extend(ack(0x50, *id)); + } + wire.send(&burst).await; + let first_rels = collect_acks(&mut wire, 6, 3).await; + assert_eq!(first_rels, ids, "MQTT-4.6.0-4: PUBREL order differs from PUBREC order"); + wire.send(&ack(0x70, ids[1])).await; + tokio::time::sleep(Duration::from_millis(200)).await; + drop(wire); + let mut wire = accept(&listener, true).await; + let (packets, _) = wire.collect(Duration::from_secs(3)).await; + let rels: Vec = packets.iter().filter(|p| p.kind() == 6).map(Raw::ack_id).collect(); + assert!( + packets.iter().all(|p| p.kind() != 3), + "MQTT-4.3.3-6: PUBLISH re-sent after PUBREL: {:?}", + descriptions(&packets) + ); + assert_eq!( + rels, + vec![ids[0], ids[2]], + "MQTT-4.4.0-1 / 4.6.0-4: outstanding PUBRELs must be re-sent in PUBREC order after Session Present=1; wire: {:?}", + descriptions(&packets) + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_3_2_2_5_session_present_zero_discards_session_state() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, persistent_options("c-sp0")).await; + let hits = Arc::new(AtomicU32::new(0)); + let counter = Arc::clone(&hits); + let (sub, ()) = tokio::join!( + client.subscribe_with_options("in", qos_options(QoS::ExactlyOnce), move |_m| { + counter.fetch_add(1, Ordering::SeqCst); + }), + async { + let _ = wire.next(Duration::from_millis(300)).await; + } + ); + sub.expect("subscribe"); + let c = client.clone(); + tokio::spawn(async move { c.publish_qos("t", b"q2".to_vec(), QoS::ExactlyOnce).await }); + let out_id = wire.expect("PUBLISH").await.publish_id().unwrap(); + wire.send(&ack(0x50, out_id)).await; + assert_eq!(wire.expect("PUBREL").await.kind(), 6); + wire.send(&server_publish(2, false, 9, "in", b"old")).await; + assert_eq!(wire.expect("PUBREC").await.kind(), 5); + assert!(wait_until(|| hits.load(Ordering::SeqCst) == 1).await); + drop(wire); + + let mut wire = accept(&listener, false).await; + let (packets, _) = wire.collect(Duration::from_secs(2)).await; + assert!( + packets.iter().all(|p| p.kind() != 3 && p.kind() != 6), + "MQTT-3.2.2-5: old session packets re-sent after Session Present=0: {:?}", + descriptions(&packets) + ); + wire.send(&server_publish(2, false, 9, "in", b"new")).await; + assert_eq!(wire.expect("PUBREC").await.kind(), 5); + assert!( + wait_until(|| hits.load(Ordering::SeqCst) == 2).await, + "MQTT-3.2.2-5: inbound QoS2 state not discarded; new PUBLISH id 9 suppressed as duplicate" + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_3_2_2_queued_while_offline_first_send_has_dup_zero() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, wire) = connected(&listener, persistent_options("c-queue")).await; + client.set_queue_on_disconnect(true).await; + drop(wire); + let probe = client.clone(); + let mut offline = false; + for _ in 0..100 { + if !probe.is_connected().await { + offline = true; + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert!(offline, "client did not notice the dropped connection"); + client + .publish_qos("t", b"queued".to_vec(), QoS::AtLeastOnce) + .await + .expect("queued publish accepted"); + let mut wire = accept(&listener, true).await; + let p = wire.expect("queued PUBLISH").await; + assert_eq!(p.kind(), 3, "expected PUBLISH, got {}", p.describe()); + assert!( + !p.publish_dup(), + "MQTT-4.3.2-2 violated: first transmission of a queued QoS1 message carries DUP=1: {}", + p.describe() + ); + })) + .await; +} + +#[tokio::test] +async fn mqtt_4_3_2_1_packet_id_not_reused_while_pubrel_outstanding() { + with_timeout(Box::pin(async { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let (client, mut wire) = connected(&listener, base_options("c-pid-wrap")).await; + let c = client.clone(); + tokio::spawn(async move { c.publish_qos("t", b"held".to_vec(), QoS::ExactlyOnce).await }); + let held = wire.expect("QoS2 PUBLISH").await.publish_id().unwrap(); + wire.send(&ack(0x50, held)).await; + assert_eq!(wire.expect("PUBREL").await.kind(), 6); + let c = client.clone(); + let publisher = tokio::spawn(async move { + for _ in 0..65_535u32 { + if c.publish_qos("t", b"x".to_vec(), QoS::AtLeastOnce).await.is_err() { + break; + } + } + }); + let mut reused = None; + for _ in 0..65_535u32 { + let p = wire.expect("QoS1 PUBLISH").await; + let id = p.publish_id().unwrap(); + if id == held { + reused = Some(id); + break; + } + wire.send(&ack(0x40, id)).await; + } + publisher.abort(); + assert!( + reused.is_none(), + "MQTT-4.3.2-1 violated: packet id {held} reassigned to a new QoS1 message while its QoS2 PUBREL is still unacknowledged" + ); + })) + .await; +} diff --git a/crates/mqtt5/tests/conf_client_d.rs b/crates/mqtt5/tests/conf_client_d.rs new file mode 100644 index 00000000..b4438c9e --- /dev/null +++ b/crates/mqtt5/tests/conf_client_d.rs @@ -0,0 +1,1735 @@ +use mqtt5::{AuthHandler, AuthResponse, ConnectOptions, MqttClient}; +use std::future::Future; +use std::net::SocketAddr; +use std::pin::Pin; +use std::sync::{Arc, Mutex}; +use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; +use tokio::time::{timeout, Instant}; + +const CONNECT: u8 = 1; +const PUBLISH: u8 = 3; +const PUBACK: u8 = 4; +const PUBREC: u8 = 5; +const PUBCOMP: u8 = 7; +const SUBSCRIBE: u8 = 8; +const UNSUBSCRIBE: u8 = 10; +const PINGREQ: u8 = 12; +const AUTH: u8 = 15; + +const PROP_SERVER_KEEP_ALIVE: u8 = 0x13; +const PROP_AUTH_METHOD: u8 = 0x15; +const PROP_AUTH_DATA: u8 = 0x16; +const PROP_RECEIVE_MAXIMUM: u8 = 0x21; + +#[derive(Debug, Clone)] +struct Pkt { + kind: u8, + flags: u8, + body: Vec, +} + +enum Read { + Packet(Pkt), + Closed, + TimedOut, +} + +fn varint(mut n: usize) -> Vec { + let mut out = Vec::new(); + loop { + let mut byte = u8::try_from(n % 128).unwrap_or(0); + n /= 128; + if n > 0 { + byte |= 0x80; + } + out.push(byte); + if n == 0 { + return out; + } + } +} + +fn frame(first: u8, body: &[u8]) -> Vec { + let mut out = vec![first]; + out.extend(varint(body.len())); + out.extend_from_slice(body); + out +} + +fn mqtt_str(s: &[u8]) -> Vec { + let len = u16::try_from(s.len()).unwrap_or(u16::MAX); + let mut out = len.to_be_bytes().to_vec(); + out.extend_from_slice(s); + out +} + +fn prop_u16(id: u8, v: u16) -> Vec { + let mut out = vec![id]; + out.extend(v.to_be_bytes()); + out +} + +fn prop_str(id: u8, s: &str) -> Vec { + let mut out = vec![id]; + out.extend(mqtt_str(s.as_bytes())); + out +} + +fn with_props(fixed: &[u8], props: &[u8]) -> Vec { + let mut body = fixed.to_vec(); + body.extend(varint(props.len())); + body.extend_from_slice(props); + body +} + +fn connack(session_present: bool, reason: u8, props: &[u8]) -> Vec { + frame( + 0x20, + &with_props(&[u8::from(session_present), reason], props), + ) +} + +fn auth(reason: u8, props: &[u8]) -> Vec { + frame(0xF0, &with_props(&[reason], props)) +} + +fn disconnect(reason: u8) -> Vec { + frame(0xE0, &[reason, 0]) +} + +fn puback(pid: u16) -> Vec { + frame(0x40, &pid.to_be_bytes()) +} + +fn pubrel(pid: u16) -> Vec { + frame(0x62, &pid.to_be_bytes()) +} + +fn publish(topic: &str, qos: u8, pid: u16, payload: &[u8]) -> Vec { + let mut body = mqtt_str(topic.as_bytes()); + if qos > 0 { + body.extend(pid.to_be_bytes()); + } + body.push(0); + body.extend_from_slice(payload); + frame(0x30 | (qos << 1), &body) +} + +fn suback(pid: u16, count: usize) -> Vec { + let mut body = pid.to_be_bytes().to_vec(); + body.push(0); + body.extend(std::iter::repeat_n(0u8, count)); + frame(0x90, &body) +} + +fn unsuback(pid: u16, count: usize) -> Vec { + let mut body = pid.to_be_bytes().to_vec(); + body.push(0); + body.extend(std::iter::repeat_n(0u8, count)); + frame(0xB0, &body) +} + +fn take_varint(b: &[u8], pos: &mut usize) -> Option { + let mut value = 0usize; + let mut mult = 1usize; + loop { + let byte = *b.get(*pos)?; + *pos += 1; + value += usize::from(byte & 0x7f) * mult; + if byte & 0x80 == 0 { + return Some(value); + } + mult *= 128; + if mult > 128 * 128 * 128 { + return None; + } + } +} + +fn take_str(b: &[u8], pos: &mut usize) -> Option> { + let len = usize::from(u16::from_be_bytes([*b.get(*pos)?, *b.get(*pos + 1)?])); + let start = *pos + 2; + let out = b.get(start..start + len)?.to_vec(); + *pos = start + len; + Some(out) +} + +fn take_u16(b: &[u8], pos: &mut usize) -> Option { + let v = u16::from_be_bytes([*b.get(*pos)?, *b.get(*pos + 1)?]); + *pos += 2; + Some(v) +} + +fn parse_props(b: &[u8], pos: &mut usize) -> Option)>> { + let len = take_varint(b, pos)?; + let end = *pos + len; + let mut out = Vec::new(); + while *pos < end { + let id = *b.get(*pos)?; + *pos += 1; + let value = match id { + 0x01 | 0x17 | 0x19 | 0x24 | 0x25 | 0x28 | 0x29 | 0x2A => { + let v = vec![*b.get(*pos)?]; + *pos += 1; + v + } + 0x13 | 0x21 | 0x22 | 0x23 => { + let v = b.get(*pos..*pos + 2)?.to_vec(); + *pos += 2; + v + } + 0x02 | 0x11 | 0x18 | 0x27 => { + let v = b.get(*pos..*pos + 4)?.to_vec(); + *pos += 4; + v + } + 0x0B => { + let v = take_varint(b, pos)?; + v.to_be_bytes().to_vec() + } + 0x03 | 0x08 | 0x12 | 0x15 | 0x1A | 0x1C | 0x1F | 0x09 | 0x16 => take_str(b, pos)?, + 0x26 => { + let mut k = take_str(b, pos)?; + k.push(b'='); + k.extend(take_str(b, pos)?); + k + } + _ => return None, + }; + out.push((id, value)); + } + Some(out) +} + +fn prop_string(props: &[(u8, Vec)], id: u8) -> Option { + props + .iter() + .find(|(pid, _)| *pid == id) + .map(|(_, v)| String::from_utf8_lossy(v).into_owned()) +} + +struct ConnectInfo { + clean_start: bool, + auth_method: Option, +} + +fn parse_connect(p: &Pkt) -> Option { + let b = &p.body; + let mut pos = 0; + take_str(b, &mut pos)?; + pos += 1; + let flags = *b.get(pos)?; + pos += 1; + take_u16(b, &mut pos)?; + let props = parse_props(b, &mut pos)?; + Some(ConnectInfo { + clean_start: flags & 0x02 != 0, + auth_method: prop_string(&props, PROP_AUTH_METHOD), + }) +} + +struct PublishInfo { + topic: String, + qos: u8, + dup: bool, + pid: Option, +} + +fn parse_publish(p: &Pkt) -> Option { + let b = &p.body; + let mut pos = 0; + let topic = String::from_utf8_lossy(&take_str(b, &mut pos)?).into_owned(); + let qos = (p.flags >> 1) & 0x03; + let pid = if qos > 0 { + Some(take_u16(b, &mut pos)?) + } else { + None + }; + Some(PublishInfo { + topic, + qos, + dup: p.flags & 0x08 != 0, + pid, + }) +} + +struct AuthInfo { + reason: u8, + method: Option, +} + +fn parse_auth(p: &Pkt) -> AuthInfo { + if p.body.is_empty() { + return AuthInfo { + reason: 0, + method: None, + }; + } + let mut pos = 1; + let method = + parse_props(&p.body, &mut pos).and_then(|props| prop_string(&props, PROP_AUTH_METHOD)); + AuthInfo { + reason: p.body[0], + method, + } +} + +fn parse_filters(p: &Pkt, with_options: bool) -> Option<(u16, Vec)> { + let b = &p.body; + let mut pos = 0; + let pid = take_u16(b, &mut pos)?; + parse_props(b, &mut pos)?; + let mut filters = Vec::new(); + while pos < b.len() { + filters.push(String::from_utf8_lossy(&take_str(b, &mut pos)?).into_owned()); + if with_options { + pos += 1; + } + } + Some((pid, filters)) +} + +fn packet_id(p: &Pkt) -> Option { + let mut pos = 0; + take_u16(&p.body, &mut pos) +} + +fn try_split_packet(buf: &mut Vec) -> Option { + let first = *buf.first()?; + let mut pos = 1; + let len = take_varint(buf, &mut pos)?; + if buf.len() < pos + len { + return None; + } + let body = buf[pos..pos + len].to_vec(); + buf.drain(..pos + len); + Some(Pkt { + kind: first >> 4, + flags: first & 0x0f, + body, + }) +} + +struct Peer { + stream: TcpStream, + buf: Vec, +} + +impl Peer { + fn new(stream: TcpStream) -> Self { + Self { + stream, + buf: Vec::new(), + } + } + + async fn read(&mut self, wait: Duration) -> Read { + let deadline = Instant::now() + wait; + loop { + if let Some(p) = try_split_packet(&mut self.buf) { + return Read::Packet(p); + } + let mut chunk = [0u8; 4096]; + match tokio::time::timeout_at(deadline, self.stream.read(&mut chunk)).await { + Err(_) => return Read::TimedOut, + Ok(Ok(0) | Err(_)) => return Read::Closed, + Ok(Ok(n)) => self.buf.extend_from_slice(&chunk[..n]), + } + } + } + + async fn expect(&mut self, wait: Duration) -> Option { + match self.read(wait).await { + Read::Packet(p) => Some(p), + Read::Closed | Read::TimedOut => None, + } + } + + async fn send(&mut self, bytes: &[u8]) { + let _ = self.stream.write_all(bytes).await; + let _ = self.stream.flush().await; + } + + async fn closed_within(&mut self, wait: Duration) -> (bool, Vec) { + let deadline = Instant::now() + wait; + let mut seen = Vec::new(); + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + match self.read(remaining).await { + Read::Packet(p) => seen.push(p.kind), + Read::Closed => return (true, seen), + Read::TimedOut => return (false, seen), + } + } + } +} + +async fn listener() -> (TcpListener, SocketAddr) { + let l = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let a = l.local_addr().expect("addr"); + (l, a) +} + +async fn accept_connect(l: &TcpListener, wait: Duration) -> Option<(Peer, ConnectInfo)> { + let (stream, _) = timeout(wait, l.accept()).await.ok()?.ok()?; + let mut peer = Peer::new(stream); + let p = peer.expect(wait).await?; + if p.kind != CONNECT { + return None; + } + let info = parse_connect(&p)?; + Some((peer, info)) +} + +fn opts(id: &str) -> ConnectOptions { + ConnectOptions::new(id).with_automatic_reconnect(false) +} + +fn persistent_opts(id: &str) -> ConnectOptions { + ConnectOptions::new(id) + .with_clean_start(false) + .with_session_expiry_interval(3600) + .with_automatic_reconnect(true) + .with_reconnect_delay(Duration::from_millis(100), Duration::from_millis(500)) +} + +const W: Duration = Duration::from_secs(3); + +#[derive(Default)] +struct Recorded { + publishes: Vec, + subscribes: Vec, + unsubscribes: Vec, +} + +async fn recorder_server() -> (SocketAddr, Arc>) { + let (l, addr) = listener().await; + let rec = Arc::new(Mutex::new(Recorded::default())); + let rec2 = Arc::clone(&rec); + tokio::spawn(async move { + let Some((mut peer, _)) = accept_connect(&l, Duration::from_secs(10)).await else { + return; + }; + peer.send(&connack(false, 0, &[])).await; + while let Read::Packet(p) = peer.read(Duration::from_secs(60)).await { + match p.kind { + PUBLISH => { + if let Some(info) = parse_publish(&p) { + rec2.lock().expect("lock").publishes.push(info.topic); + if let (1, Some(pid)) = (info.qos, info.pid) { + peer.send(&puback(pid)).await; + } + } + } + SUBSCRIBE => { + if let Some((pid, filters)) = parse_filters(&p, true) { + let n = filters.len(); + rec2.lock().expect("lock").subscribes.extend(filters); + peer.send(&suback(pid, n)).await; + } + } + UNSUBSCRIBE => { + if let Some((pid, filters)) = parse_filters(&p, false) { + let n = filters.len(); + rec2.lock().expect("lock").unsubscribes.extend(filters); + peer.send(&unsuback(pid, n)).await; + } + } + PINGREQ => peer.send(&[0xD0, 0x00]).await, + _ => {} + } + } + }); + (addr, rec) +} + +async fn recorder_client(id: &str) -> (MqttClient, Arc>) { + let (addr, rec) = recorder_server().await; + let client = MqttClient::with_options(opts(id)); + client + .connect(&format!("mqtt://{addr}")) + .await + .expect("connect to recorder"); + (client, rec) +} + +async fn publish_cases( + client: &MqttClient, + rec: &Arc>, + topics: &[String], +) -> Vec { + for t in topics { + let _ = timeout(W, client.publish(t.clone(), b"x".to_vec())).await; + } + tokio::time::sleep(Duration::from_millis(300)).await; + let r = rec.lock().expect("lock"); + topics + .iter() + .filter(|t| r.publishes.contains(t)) + .map(|t| format!("{:?}", preview(t))) + .collect() +} + +async fn subscribe_cases( + client: &MqttClient, + rec: &Arc>, + filters: &[String], +) -> Vec { + for f in filters { + let _ = timeout(W, client.subscribe(f.clone(), |_| {})).await; + let _ = timeout(W, client.unsubscribe(f.clone())).await; + } + tokio::time::sleep(Duration::from_millis(300)).await; + let r = rec.lock().expect("lock"); + filters + .iter() + .filter(|f| r.subscribes.contains(f) || r.unsubscribes.contains(f)) + .map(|f| { + let which = match (r.subscribes.contains(f), r.unsubscribes.contains(f)) { + (true, true) => "SUBSCRIBE+UNSUBSCRIBE", + (true, false) => "SUBSCRIBE", + _ => "UNSUBSCRIBE", + }; + format!("{:?} via {which}", preview(f)) + }) + .collect() +} + +fn preview(s: &str) -> String { + if s.len() > 40 { + format!("<{} bytes>", s.len()) + } else { + s.to_string() + } +} + +fn strings(items: &[&str]) -> Vec { + items.iter().map(|s| (*s).to_string()).collect() +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_7_0_1_client_must_not_send_wildcards_in_topic_name() { + let (client, rec) = recorder_client("d-4701").await; + let sent = publish_cases( + &client, + &rec, + &strings(&["a/+", "a/#", "+", "#", "sport+", "a/b#"]), + ) + .await; + assert!( + sent.is_empty(), + "MQTT-4.7.0-1 VIOLATION: client put wildcard Topic Names on the wire in PUBLISH: {sent:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_7_1_1_client_must_not_send_misplaced_multilevel_wildcard() { + let (client, rec) = recorder_client("d-4711").await; + let sent = subscribe_cases( + &client, + &rec, + &strings(&["a/#/b", "a#", "#/a", "a/b#", "##"]), + ) + .await; + assert!( + sent.is_empty(), + "MQTT-4.7.1-1 VIOLATION: client sent Topic Filters with misplaced '#': {sent:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_7_1_2_client_must_not_send_partial_level_single_wildcard() { + let (client, rec) = recorder_client("d-4712").await; + let sent = subscribe_cases( + &client, + &rec, + &strings(&["a+", "a/+b", "+a/b", "a/b+/c", "++"]), + ) + .await; + assert!( + sent.is_empty(), + "MQTT-4.7.1-2 VIOLATION: client sent Topic Filters where '+' does not occupy a whole level: {sent:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_7_3_1_client_must_not_send_empty_topic_or_filter() { + let (client, rec) = recorder_client("d-4731").await; + let empty = vec![String::new()]; + let published = publish_cases(&client, &rec, &empty).await; + let subscribed = subscribe_cases(&client, &rec, &empty).await; + assert!( + published.is_empty() && subscribed.is_empty(), + "MQTT-4.7.3-1 VIOLATION: client sent zero-length Topic Name (PUBLISH, no Topic Alias) {published:?} / Topic Filter {subscribed:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_7_3_2_client_must_not_send_null_character() { + let (client, rec) = recorder_client("d-4732").await; + let cases = strings(&["a\0b", "\0"]); + let published = publish_cases(&client, &rec, &cases).await; + let subscribed = subscribe_cases(&client, &rec, &cases).await; + assert!( + published.is_empty() && subscribed.is_empty(), + "MQTT-4.7.3-2 VIOLATION: client sent U+0000 in Topic Name {published:?} / Topic Filter {subscribed:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_7_3_3_client_must_not_send_topic_over_65535_bytes() { + let (client, rec) = recorder_client("d-4733").await; + let cases = vec!["a".repeat(65_536), "é".repeat(32_768)]; + let published = publish_cases(&client, &rec, &cases).await; + let subscribed = subscribe_cases(&client, &rec, &cases).await; + assert!( + published.is_empty() && subscribed.is_empty(), + "MQTT-4.7.3-3 VIOLATION: client sent >65535-byte Topic Name {published:?} / Topic Filter {subscribed:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_8_2_1_client_must_not_send_share_without_sharename_or_filter() { + let (client, rec) = recorder_client("d-4821").await; + let sent = subscribe_cases( + &client, + &rec, + &strings(&[ + "$share//x", + "$share/g", + "$share/g/", + "$share/", + "$share/g/a/#/b", + ]), + ) + .await; + assert!( + sent.is_empty(), + "MQTT-4.8.2-1 VIOLATION: client sent malformed Shared Subscription filters (empty ShareName / missing or invalid filter): {sent:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_8_2_2_client_must_not_send_wildcard_in_sharename() { + let (client, rec) = recorder_client("d-4822").await; + let sent = subscribe_cases( + &client, + &rec, + &strings(&["$share/g+/x", "$share/g#/x", "$share/+/x", "$share/#/x"]), + ) + .await; + assert!( + sent.is_empty(), + "MQTT-4.8.2-2 VIOLATION: client sent Shared Subscription ShareName containing '+' or '#': {sent:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_7_and_4_8_client_must_not_reject_valid_filters() { + let (client, rec) = recorder_client("d-47valid").await; + let valid = strings(&[ + "$shared/x", + "$sharex", + "$share/g/a/+", + "$share/g/#", + "/", + "+/+", + "a/+/b/#", + "+", + "#", + "$SYS/#", + "a//b", + ]); + let mut outcomes = Vec::new(); + for f in &valid { + let r = timeout(W, client.subscribe(f.clone(), |_| {})).await; + outcomes.push((f.clone(), matches!(r, Ok(Ok(_))))); + } + tokio::time::sleep(Duration::from_millis(300)).await; + let r = rec.lock().expect("lock"); + let rejected: Vec<_> = outcomes + .iter() + .filter(|(f, ok)| !ok || !r.subscribes.contains(f)) + .map(|(f, _)| f.clone()) + .collect(); + assert!( + rejected.is_empty(), + "MQTT-4.7.1/4.8.2 over-rejection: client refused valid Topic Filters: {rejected:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_9_0_1_initial_send_quota_is_receive_maximum() { + let (l, addr) = listener().await; + let seen = Arc::new(Mutex::new(0usize)); + let seen2 = Arc::clone(&seen); + tokio::spawn(async move { + let Some((mut peer, _)) = accept_connect(&l, W).await else { + return; + }; + peer.send(&connack(false, 0, &prop_u16(PROP_RECEIVE_MAXIMUM, 2))) + .await; + while let Read::Packet(p) = peer.read(Duration::from_secs(30)).await { + if p.kind == PUBLISH { + *seen2.lock().expect("lock") += 1; + } + } + }); + let client = MqttClient::with_options(opts("d-4901")); + client + .connect(&format!("mqtt://{addr}")) + .await + .expect("connect"); + let mut handles = Vec::new(); + for i in 0..4u8 { + let c = client.clone(); + handles.push(tokio::spawn( + async move { c.publish_qos1("q/t", vec![i]).await }, + )); + } + tokio::time::sleep(Duration::from_millis(700)).await; + let n = *seen.lock().expect("lock"); + for h in handles { + h.abort(); + } + assert_eq!(n, 2, "MQTT-4.9.0-1: initial send quota must equal Receive Maximum 2, saw {n} unacked QoS1 PUBLISH"); +} + +struct QuotaTrace { + distinct_unacked: usize, + dup_resends: usize, +} + +async fn count_unacked_publishes(peer: &mut Peer, window: Duration) -> (QuotaTrace, Vec) { + let deadline = Instant::now() + window; + let mut pids = Vec::new(); + let mut dup = 0; + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + match peer.read(remaining).await { + Read::Packet(p) if p.kind == PUBLISH => { + if let Some(info) = parse_publish(&p) { + if info.dup { + dup += 1; + } + if let (true, Some(pid)) = (info.qos > 0, info.pid) { + if !pids.contains(&pid) { + pids.push(pid); + } + } + } + } + Read::Packet(p) if p.kind == PINGREQ => peer.send(&[0xD0, 0x00]).await, + Read::Packet(_) => {} + Read::Closed | Read::TimedOut => break, + } + } + ( + QuotaTrace { + distinct_unacked: pids.len(), + dup_resends: dup, + }, + pids, + ) +} + +async fn wait_disconnected(client: &MqttClient) -> bool { + let deadline = Instant::now() + W; + while Instant::now() < deadline { + if let Ok(false) = timeout(Duration::from_millis(200), client.is_connected()).await { + return true; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + false +} + +async fn wait_connected(client: &MqttClient) -> bool { + let deadline = Instant::now() + Duration::from_secs(8); + while Instant::now() < deadline { + if let Ok(true) = timeout(Duration::from_millis(200), client.is_connected()).await { + return true; + } + tokio::time::sleep(Duration::from_millis(20)).await; + } + false +} + +async fn first_connection_with_two_unacked( + l: &TcpListener, + addr: SocketAddr, + id: &str, +) -> ( + MqttClient, + Vec>>, +) { + let client = MqttClient::with_options(persistent_opts(id)); + let connect = { + let c = client.clone(); + tokio::spawn(async move { c.connect(&format!("mqtt://{addr}")).await }) + }; + let (mut peer, _) = accept_connect(l, W).await.expect("conn1 CONNECT"); + peer.send(&connack(false, 0, &[])).await; + connect.await.expect("join").expect("conn1 connect"); + let mut handles = Vec::new(); + for i in 0..2u8 { + let c = client.clone(); + handles.push(tokio::spawn(async move { + c.publish_qos1("q/resume", vec![i]).await + })); + } + let (trace, _) = count_unacked_publishes(&mut peer, Duration::from_millis(500)).await; + assert_eq!( + trace.distinct_unacked, 2, + "setup: conn1 must carry 2 unacked QoS1 PUBLISH" + ); + drop(peer); + assert!( + wait_disconnected(&client).await, + "setup: client must notice conn1 loss" + ); + (client, handles) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_9_0_2_send_quota_after_resume_with_smaller_receive_maximum() { + let (l, addr) = listener().await; + let (client, first_handles) = + first_connection_with_two_unacked(&l, addr, "d-4902-resume").await; + + let (mut peer, info) = accept_connect(&l, Duration::from_secs(8)) + .await + .expect("conn2 CONNECT"); + assert!( + !info.clean_start, + "setup: reconnect must resume the session (Clean Start 0)" + ); + peer.send(&connack(true, 0, &prop_u16(PROP_RECEIVE_MAXIMUM, 1))) + .await; + assert!( + wait_connected(&client).await, + "setup: client must be connected on conn2" + ); + + let mut later = Vec::new(); + for i in 0..2u8 { + let c = client.clone(); + later.push(tokio::spawn(async move { + c.publish_qos1("q/resume", vec![10 + i]).await + })); + } + let (trace, pids) = count_unacked_publishes(&mut peer, Duration::from_millis(1500)).await; + eprintln!( + "MQTT-4.9.0-2 resume trace: unacked QoS1 on conn2 = {} (pids {pids:?}), DUP resends = {}", + trace.distinct_unacked, trace.dup_resends + ); + + for pid in &pids { + peer.send(&puback(*pid)).await; + } + let server = tokio::spawn(async move { + while let Read::Packet(p) = peer.read(Duration::from_secs(20)).await { + match p.kind { + PUBLISH => { + if let Some(pid) = parse_publish(&p).and_then(|i| i.pid) { + peer.send(&puback(pid)).await; + } + } + PINGREQ => peer.send(&[0xD0, 0x00]).await, + _ => {} + } + } + }); + + let mut panicked = 0; + let mut hung = 0; + for h in first_handles.into_iter().chain(later) { + match timeout(Duration::from_secs(15), h).await { + Ok(Ok(_)) => {} + Ok(Err(_)) => panicked += 1, + Err(_) => hung += 1, + } + } + let fresh = timeout( + Duration::from_secs(5), + client.publish_qos1("q/resume", b"fresh".to_vec()), + ) + .await; + server.abort(); + + assert_eq!( + panicked, 0, + "MQTT-4.9.0-2: publish task panicked after resume" + ); + assert_eq!( + hung, 0, + "MQTT-4.9.0-2: {hung} publish calls deadlocked after resume with smaller Receive Maximum" + ); + assert!( + trace.distinct_unacked <= 1, + "MQTT-4.9.0-2 VIOLATION: after resume with Receive Maximum 1 the client had {} unacked QoS>0 PUBLISH on the wire ({} DUP)", + trace.distinct_unacked, + trace.dup_resends + ); + assert!( + matches!(fresh, Ok(Ok(_))), + "MQTT-4.9.0-2: send quota leaked after resume; a fresh acked QoS1 publish failed: {fresh:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_9_0_2_offline_queue_flush_respects_receive_maximum() { + let (l, addr) = listener().await; + let client = MqttClient::with_options( + persistent_opts("d-4902-queue") + .with_reconnect_delay(Duration::from_millis(600), Duration::from_secs(1)), + ); + let connect = { + let c = client.clone(); + tokio::spawn(async move { c.connect(&format!("mqtt://{addr}")).await }) + }; + let (peer, _) = accept_connect(&l, W).await.expect("conn1 CONNECT"); + let mut peer = peer; + peer.send(&connack(false, 0, &[])).await; + connect.await.expect("join").expect("conn1 connect"); + drop(peer); + assert!( + wait_disconnected(&client).await, + "setup: client must notice conn1 loss" + ); + + let mut queued = 0; + for i in 0..3u8 { + if let Ok(Ok(_)) = timeout( + Duration::from_millis(300), + client.publish_qos1("q/offline", vec![i]), + ) + .await + { + queued += 1; + } + } + assert_eq!( + queued, 3, + "setup: 3 QoS1 publishes must be accepted into the offline queue" + ); + + let (mut peer, _) = accept_connect(&l, Duration::from_secs(8)) + .await + .expect("conn2 CONNECT"); + peer.send(&connack(true, 0, &prop_u16(PROP_RECEIVE_MAXIMUM, 1))) + .await; + let (trace, _) = count_unacked_publishes(&mut peer, Duration::from_millis(1500)).await; + assert!( + trace.distinct_unacked <= 1, + "MQTT-4.9.0-2 VIOLATION: offline-queue flush after reconnect sent {} unacked QoS1 PUBLISH with server Receive Maximum 1 (send quota bypassed)", + trace.distinct_unacked + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_9_0_1_send_quota_reinitialized_after_ack_timeout_and_reconnect() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("d-4901-reinit")); + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("conn1 CONNECT"); + peer.send(&connack(false, 0, &prop_u16(PROP_RECEIVE_MAXIMUM, 1))) + .await; + connect.await.expect("join").expect("conn1 connect"); + + let timed_out = timeout( + Duration::from_secs(15), + client.publish_qos1("q/reinit", b"unacked".to_vec()), + ) + .await; + assert!( + matches!(timed_out, Ok(Err(mqtt5::MqttError::Timeout))), + "setup: unacked QoS1 publish must hit the client ack timeout: {timed_out:?}" + ); + drop(peer); + assert!( + wait_disconnected(&client).await, + "setup: client must notice conn1 loss" + ); + + let (mut peer, _) = accept_connect(&l, Duration::from_secs(8)) + .await + .expect("conn2 CONNECT"); + peer.send(&connack(false, 0, &prop_u16(PROP_RECEIVE_MAXIMUM, 1))) + .await; + assert!(wait_connected(&client).await, "setup: reconnect"); + + let fresh = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("q/reinit", b"fresh".to_vec()).await }) + }; + let mut delivered = false; + let deadline = Instant::now() + W; + while Instant::now() < deadline { + let remaining = deadline.saturating_duration_since(Instant::now()); + match peer.read(remaining).await { + Read::Packet(p) if p.kind == PUBLISH => { + if let Some(pid) = parse_publish(&p).and_then(|i| i.pid) { + peer.send(&puback(pid)).await; + } + delivered = true; + break; + } + Read::Packet(p) if p.kind == PINGREQ => peer.send(&[0xD0, 0x00]).await, + Read::Packet(_) => {} + Read::Closed | Read::TimedOut => break, + } + } + fresh.abort(); + assert!( + delivered, + "MQTT-4.9.0-1 VIOLATION: on a new Network Connection (Session Present 0, Receive Maximum 1) the send quota was not re-initialized; the permit held by a pre-disconnect ack-timeout leaked, initial quota is 0 and a fresh QoS1 PUBLISH never reached the wire within 3s" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn out_of_group_mqtt_4_4_0_1_resend_unacked_publish_on_resume() { + let (l, addr) = listener().await; + let (client, handles) = first_connection_with_two_unacked(&l, addr, "d-4401").await; + let (mut peer, _) = accept_connect(&l, Duration::from_secs(8)) + .await + .expect("conn2 CONNECT"); + peer.send(&connack(true, 0, &[])).await; + assert!(wait_connected(&client).await, "setup: reconnect"); + let (trace, _) = count_unacked_publishes(&mut peer, Duration::from_millis(1500)).await; + for h in handles { + h.abort(); + } + assert_eq!( + trace.dup_resends, 2, + "MQTT-4.4.0-1 VIOLATION (outside group D): on Session resume the client resent {} of 2 unacknowledged QoS1 PUBLISH packets", + trace.dup_resends + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_9_0_3_client_still_acks_and_pings_at_zero_quota() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-4903").with_keep_alive(Duration::from_secs(1))); + let connect = { + let c = client.clone(); + tokio::spawn(async move { c.connect(&format!("mqtt://{addr}")).await }) + }; + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + let mut props = prop_u16(PROP_RECEIVE_MAXIMUM, 1); + props.extend(prop_u16(PROP_SERVER_KEEP_ALIVE, 1)); + peer.send(&connack(false, 0, &props)).await; + connect.await.expect("join").expect("connect"); + + let blocked = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("q/zero", b"hold".to_vec()).await }) + }; + let first = peer.expect(W).await.expect("first QoS1 PUBLISH"); + assert_eq!(first.kind, PUBLISH); + let queued_publish = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("q/zero", b"blocked".to_vec()).await }) + }; + + peer.send(&publish("in/q1", 1, 700, b"a")).await; + peer.send(&publish("in/q2", 2, 701, b"b")).await; + let subscribe = { + let c = client.clone(); + tokio::spawn(async move { c.subscribe("in/#", |_| {}).await }) + }; + + let mut got_puback = false; + let mut got_pubrec = false; + let mut got_pubcomp = false; + let mut got_ping = false; + let mut got_subscribe = false; + let mut extra_publish = 0; + let deadline = Instant::now() + Duration::from_secs(4); + while Instant::now() < deadline { + let remaining = deadline.saturating_duration_since(Instant::now()); + match peer.read(remaining).await { + Read::Packet(p) => match p.kind { + PUBACK if packet_id(&p) == Some(700) => got_puback = true, + PUBREC if packet_id(&p) == Some(701) => { + got_pubrec = true; + peer.send(&pubrel(701)).await; + } + PUBCOMP if packet_id(&p) == Some(701) => got_pubcomp = true, + PINGREQ => { + got_ping = true; + peer.send(&[0xD0, 0x00]).await; + } + SUBSCRIBE => { + got_subscribe = true; + if let Some((pid, f)) = parse_filters(&p, true) { + peer.send(&suback(pid, f.len())).await; + } + } + PUBLISH => extra_publish += 1, + _ => {} + }, + Read::Closed | Read::TimedOut => break, + } + if got_puback && got_pubrec && got_pubcomp && got_ping && got_subscribe { + break; + } + } + let sub_ok = matches!(timeout(W, subscribe).await, Ok(Ok(Ok(_)))); + blocked.abort(); + queued_publish.abort(); + assert_eq!( + extra_publish, 0, + "MQTT-4.9.0-2: client sent a QoS1 PUBLISH while quota was 0" + ); + assert!( + got_puback && got_pubrec && got_pubcomp && got_ping && got_subscribe && sub_ok, + "MQTT-4.9.0-3 VIOLATION at quota 0: PUBACK={got_puback} PUBREC={got_pubrec} PUBCOMP={got_pubcomp} PINGREQ={got_ping} SUBSCRIBE={got_subscribe} SUBACK-processed={sub_ok}" + ); +} + +struct ScriptedAuth { + calls: Arc>>, +} + +impl AuthHandler for ScriptedAuth { + fn handle_challenge<'a>( + &'a self, + auth_method: &'a str, + _challenge_data: Option<&'a [u8]>, + ) -> Pin> + Send + 'a>> { + self.calls + .lock() + .expect("lock") + .push(auth_method.to_string()); + Box::pin(async move { Ok(AuthResponse::Continue(b"resp".to_vec())) }) + } + + fn initial_response<'a>( + &'a self, + _auth_method: &'a str, + ) -> Pin>>> + Send + 'a>> { + Box::pin(async move { Ok(Some(b"init".to_vec())) }) + } +} + +fn scripted() -> (ScriptedAuth, Arc>>) { + let calls = Arc::new(Mutex::new(Vec::new())); + ( + ScriptedAuth { + calls: Arc::clone(&calls), + }, + calls, + ) +} + +fn start_connect( + client: &MqttClient, + addr: SocketAddr, +) -> tokio::task::JoinHandle> { + let c = client.clone(); + tokio::spawn(async move { c.connect(&format!("mqtt://{addr}")).await }) +} + +async fn auth_packets_until_close(peer: &mut Peer, wait: Duration) -> (Vec, bool) { + let deadline = Instant::now() + wait; + let mut auths = Vec::new(); + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + match peer.read(remaining).await { + Read::Packet(p) if p.kind == AUTH => auths.push(parse_auth(&p)), + Read::Packet(_) => {} + Read::Closed => return (auths, true), + Read::TimedOut => return (auths, false), + } + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_0_7_no_method_no_handler_auth_during_connect() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41207a")); + let connect = start_connect(&client, addr); + let (mut peer, info) = accept_connect(&l, W).await.expect("CONNECT"); + assert!(info.auth_method.is_none()); + peer.send(&auth(0x18, &prop_str(PROP_AUTH_METHOD, "X"))) + .await; + let (auths, closed) = auth_packets_until_close(&mut peer, W).await; + let result = timeout(W, connect).await; + assert!( + auths.is_empty(), + "MQTT-4.12.0-7 VIOLATION: client without Authentication Method sent AUTH" + ); + assert!( + closed && matches!(result, Ok(Ok(Err(_)))), + "MQTT-4.12.0-6 (client side): unexpected AUTH must fail the connect and close the connection; closed={closed} result={result:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_0_7_no_method_with_handler_auth_during_connect() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41207b")); + let (handler, _) = scripted(); + client.set_auth_handler(handler).await; + let connect = start_connect(&client, addr); + let (mut peer, info) = accept_connect(&l, W).await.expect("CONNECT"); + assert!( + info.auth_method.is_none(), + "setup: CONNECT carries no Authentication Method" + ); + peer.send(&auth(0x18, &prop_str(PROP_AUTH_METHOD, "X"))) + .await; + let got = peer.expect(Duration::from_secs(1)).await; + connect.abort(); + let sent_auth = got.as_ref().filter(|p| p.kind == AUTH).map(parse_auth); + assert!( + sent_auth.is_none(), + "MQTT-4.12.0-7 VIOLATION: CONNECT had no Authentication Method but client answered server AUTH with AUTH reason=0x{:02X} method={:?}", + sent_auth.as_ref().map_or(0, |a| a.reason), + sent_auth.as_ref().and_then(|a| a.method.clone()) + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_0_7_no_method_with_handler_auth_after_connack() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41207c")); + let (handler, _) = scripted(); + client.set_auth_handler(handler).await; + let connect = start_connect(&client, addr); + let (mut peer, info) = accept_connect(&l, W).await.expect("CONNECT"); + assert!(info.auth_method.is_none()); + peer.send(&connack(false, 0, &[])).await; + connect.await.expect("join").expect("connect"); + peer.send(&auth(0x18, &prop_str(PROP_AUTH_METHOD, "X"))) + .await; + let got = peer.expect(Duration::from_secs(1)).await; + let sent_auth = got.as_ref().filter(|p| p.kind == AUTH).map(parse_auth); + assert!( + sent_auth.is_none(), + "MQTT-4.12.0-7 VIOLATION: after CONNACK, client without Authentication Method answered server AUTH with AUTH reason=0x{:02X} method={:?}", + sent_auth.as_ref().map_or(0, |a| a.reason), + sent_auth.as_ref().and_then(|a| a.method.clone()) + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_0_3_and_4_12_0_5_continue_auth_uses_0x18_and_connect_method() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41203").with_authentication_method("M-1")); + let (handler, _) = scripted(); + client.set_auth_handler(handler).await; + let connect = start_connect(&client, addr); + let (mut peer, info) = accept_connect(&l, W).await.expect("CONNECT"); + assert_eq!(info.auth_method.as_deref(), Some("M-1")); + peer.send(&auth(0x18, &prop_str(PROP_AUTH_METHOD, "M-1"))) + .await; + let reply = peer.expect(W).await.expect("client AUTH"); + let a = parse_auth(&reply); + peer.send(&connack(false, 0, &prop_str(PROP_AUTH_METHOD, "M-1"))) + .await; + let result = timeout(W, connect).await; + assert_eq!(reply.kind, AUTH, "client must answer with AUTH"); + assert_eq!( + a.reason, 0x18, + "MQTT-4.12.0-3 VIOLATION: client continuation AUTH reason 0x{:02X}", + a.reason + ); + assert_eq!( + a.method.as_deref(), + Some("M-1"), + "MQTT-4.12.0-5 VIOLATION: client AUTH method differs from CONNECT" + ); + assert!( + matches!(result, Ok(Ok(Ok(())))), + "enhanced auth connect must succeed: {result:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_0_5_server_auth_with_different_method() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41205").with_authentication_method("M-1")); + let (handler, calls) = scripted(); + client.set_auth_handler(handler).await; + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + peer.send(&auth(0x18, &prop_str(PROP_AUTH_METHOD, "OTHER"))) + .await; + let (auths, closed) = auth_packets_until_close(&mut peer, Duration::from_millis(1500)).await; + connect.abort(); + let wrong: Vec<_> = auths + .iter() + .filter(|a| a.method.as_deref() != Some("M-1")) + .collect(); + assert!( + wrong.is_empty(), + "MQTT-4.12.0-5 VIOLATION: client sent AUTH with a method other than its CONNECT method" + ); + eprintln!( + "MQTT-4.12.0-5 observation: server AUTH method mismatch -> client replied with {} AUTH(s) using CONNECT method, handler saw {:?}, connection closed by client={closed}", + auths.len(), + calls.lock().expect("lock") + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_0_2_non_continue_reason_during_connect_is_rejected() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41202").with_authentication_method("M-1")); + let (handler, _) = scripted(); + client.set_auth_handler(handler).await; + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + peer.send(&auth(0x19, &prop_str(PROP_AUTH_METHOD, "M-1"))) + .await; + let (auths, closed) = auth_packets_until_close(&mut peer, W).await; + let result = timeout(W, connect).await; + assert!( + auths.is_empty() && closed && matches!(result, Ok(Ok(Err(_)))), + "MQTT-4.12.0-2 (client side): AUTH 0x19 during CONNECT must be rejected and the connection closed; client AUTHs={} closed={closed} result={result:?}", + auths.len() + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_0_1_and_4_12_0_4_connack_failure_closes_connection() { + for reason in [0x8Cu8, 0x87] { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41204").with_authentication_method("M-1")); + let (handler, _) = scripted(); + client.set_auth_handler(handler).await; + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + peer.send(&auth(0x18, &prop_str(PROP_AUTH_METHOD, "M-1"))) + .await; + let _ = peer.expect(W).await; + peer.send(&connack(false, reason, &[])).await; + let (closed, _) = peer.closed_within(W).await; + let result = timeout(W, connect).await; + assert!( + closed && matches!(result, Ok(Ok(Err(_)))), + "MQTT-4.12.0-4/4.12.0-1: CONNACK 0x{reason:02X} mid-auth must fail connect and close; closed={closed} result={result:?}" + ); + } +} + +async fn authenticated_session(id: &str, keep_alive: Duration) -> (MqttClient, Peer) { + let (l, addr) = listener().await; + let client = MqttClient::with_options( + opts(id) + .with_authentication_method("M-1") + .with_keep_alive(keep_alive), + ); + let (handler, _) = scripted(); + client.set_auth_handler(handler).await; + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + peer.send(&connack(false, 0, &prop_str(PROP_AUTH_METHOD, "M-1"))) + .await; + connect.await.expect("join").expect("connect"); + (client, peer) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_1_1_reauth_uses_0x19_and_original_method() { + let (client, mut peer) = authenticated_session("d-41211", Duration::from_secs(60)).await; + client.reauthenticate().await.expect("reauthenticate"); + let p = peer.expect(W).await.expect("re-auth AUTH"); + let a = parse_auth(&p); + assert_eq!(p.kind, AUTH); + assert_eq!( + a.reason, 0x19, + "MQTT-4.12.1-1: re-auth must use reason 0x19" + ); + assert_eq!( + a.method.as_deref(), + Some("M-1"), + "MQTT-4.12.1-1 VIOLATION: re-auth method differs from original" + ); + peer.send(&auth(0x18, &{ + let mut pr = prop_str(PROP_AUTH_METHOD, "M-1"); + pr.push(PROP_AUTH_DATA); + pr.extend(mqtt_str(b"chal")); + pr + })) + .await; + let p2 = peer.expect(W).await.expect("re-auth continuation"); + let a2 = parse_auth(&p2); + assert!( + p2.kind == AUTH && a2.reason == 0x18 && a2.method.as_deref() == Some("M-1"), + "MQTT-4.12.0-3/4.12.0-5 during re-auth: continuation reason=0x{:02X} method={:?}", + a2.reason, + a2.method + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_1_2_server_disconnect_0x87_during_reauth_closes_connection() { + let (client, mut peer) = authenticated_session("d-41212", Duration::from_secs(1)).await; + client.reauthenticate().await.expect("reauthenticate"); + let p = peer.expect(W).await.expect("re-auth AUTH"); + assert_eq!(p.kind, AUTH); + peer.send(&disconnect(0x87)).await; + let (closed, after) = peer.closed_within(Duration::from_secs(4)).await; + let pinged = after.contains(&PINGREQ); + assert!( + closed, + "MQTT-4.12.1-2 / MQTT-4.13.2-1 VIOLATION: after server DISCONNECT 0x87 during re-authentication the client kept the TCP connection open for 4s (reconnect disabled); PINGREQ sent afterwards={pinged}, packet types seen={after:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_12_1_2_reauth_failure_auth_reason_closes_connection() { + let (client, mut peer) = authenticated_session("d-41212b", Duration::from_secs(60)).await; + client.reauthenticate().await.expect("reauthenticate"); + let _ = peer.expect(W).await; + peer.send(&auth(0x87, &[])).await; + let (closed, after) = peer.closed_within(Duration::from_secs(4)).await; + assert!( + closed, + "MQTT-4.12.1-2 VIOLATION: client treated failed re-authentication (AUTH 0x87) as fatal but never closed the Network Connection; packets seen={after:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_13_2_1_connack_error_closes_connection() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41321a")); + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + peer.send(&connack(false, 0x87, &[])).await; + let (closed, _) = peer.closed_within(W).await; + let result = timeout(W, connect).await; + assert!( + closed && matches!(result, Ok(Ok(Err(_)))), + "MQTT-4.13.2-1: CONNACK 0x87 must fail connect and close; closed={closed} result={result:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_13_2_1_server_disconnect_error_closes_connection() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(opts("d-41321b").with_keep_alive(Duration::from_secs(1))); + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + peer.send(&connack(false, 0, &[])).await; + connect.await.expect("join").expect("connect"); + peer.send(&disconnect(0x8E)).await; + let (closed, after) = peer.closed_within(Duration::from_secs(4)).await; + let connected = client.is_connected().await; + assert!( + closed, + "MQTT-4.13.2-1 VIOLATION: after server DISCONNECT 0x8E the client did not close the Network Connection within 4s (is_connected={connected}, reconnect disabled); packets sent afterwards={after:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn mqtt_4_13_2_1_server_disconnect_error_closes_connection_default_reconnect() { + let (l, addr) = listener().await; + let client = MqttClient::with_options( + ConnectOptions::new("d-41321c").with_keep_alive(Duration::from_secs(1)), + ); + let connect = start_connect(&client, addr); + let (mut peer, _) = accept_connect(&l, W).await.expect("CONNECT"); + peer.send(&connack(false, 0, &[])).await; + connect.await.expect("join").expect("connect"); + drop(l); + let t0 = Instant::now(); + peer.send(&disconnect(0x8E)).await; + let (closed, after) = peer.closed_within(Duration::from_secs(6)).await; + let elapsed = t0.elapsed(); + let _ = timeout(W, client.disconnect()).await; + assert!( + closed && elapsed < Duration::from_millis(500), + "MQTT-4.13.2-1 VIOLATION (default reconnect config): connection closed={closed} only after {elapsed:?} (tied to the reconnect attempt, not the DISCONNECT); packets sent afterwards={after:?}" + ); +} + +#[cfg(feature = "transport-websocket")] +mod websocket { + use super::{ + connack, parse_filters, publish, suback, try_split_packet, Pkt, CONNECT, SUBSCRIBE, W, + }; + use futures_util::{SinkExt, StreamExt}; + use mqtt5::{ConnectOptions, MqttClient}; + use std::net::SocketAddr; + use std::sync::{Arc, Mutex}; + use std::time::Duration; + use tokio::net::{TcpListener, TcpStream}; + use tokio::time::{timeout, Instant}; + use tokio_tungstenite::tungstenite::handshake::server::{ + Callback, ErrorResponse, Request, Response, + }; + use tokio_tungstenite::tungstenite::Message; + use tokio_tungstenite::WebSocketStream; + + struct WsPeer { + ws: WebSocketStream, + buf: Vec, + non_binary_from_client: Vec, + } + + enum WsRead { + Packet(Pkt), + Closed, + TimedOut, + } + + impl WsPeer { + async fn read(&mut self, wait: Duration) -> WsRead { + let deadline = Instant::now() + wait; + loop { + if let Some(p) = try_split_packet(&mut self.buf) { + return WsRead::Packet(p); + } + match tokio::time::timeout_at(deadline, self.ws.next()).await { + Err(_) => return WsRead::TimedOut, + Ok(None | Some(Err(_) | Ok(Message::Close(_)))) => return WsRead::Closed, + Ok(Some(Ok(Message::Binary(b)))) => self.buf.extend_from_slice(&b), + Ok(Some(Ok(other))) => self.non_binary_from_client.push(format!("{other:?}")), + } + } + } + + async fn send_binary(&mut self, bytes: Vec) { + let _ = self.ws.send(Message::Binary(bytes.into())).await; + } + + async fn closed_within(&mut self, wait: Duration) -> (bool, Vec) { + let deadline = Instant::now() + wait; + let mut seen = Vec::new(); + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + match self.read(remaining).await { + WsRead::Packet(p) => seen.push(p.kind), + WsRead::Closed => return (true, seen), + WsRead::TimedOut => return (false, seen), + } + } + } + } + + async fn ws_listener() -> (TcpListener, SocketAddr) { + let l = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let a = l.local_addr().expect("addr"); + (l, a) + } + + async fn ws_accept(l: &TcpListener) -> (WsPeer, Option) { + let (stream, _) = timeout(W, l.accept()) + .await + .expect("accept") + .expect("accept"); + let offered = Arc::new(Mutex::new(None)); + let capture = OfferedProtocols { + offered: Arc::clone(&offered), + }; + let ws = tokio_tungstenite::accept_hdr_async(stream, capture) + .await + .expect("ws handshake"); + let value = offered.lock().expect("lock").clone(); + ( + WsPeer { + ws, + buf: Vec::new(), + non_binary_from_client: Vec::new(), + }, + value, + ) + } + + struct OfferedProtocols { + offered: Arc>>, + } + + impl Callback for OfferedProtocols { + fn on_request(self, req: &Request, mut resp: Response) -> Result { + let value = req + .headers() + .get("Sec-WebSocket-Protocol") + .and_then(|v| v.to_str().ok()) + .map(str::to_string); + *self.offered.lock().expect("lock") = value; + resp.headers_mut() + .insert("Sec-WebSocket-Protocol", "mqtt".parse().expect("header")); + Ok(resp) + } + } + + fn ws_opts(id: &str) -> ConnectOptions { + ConnectOptions::new(id).with_automatic_reconnect(false) + } + + async fn ws_connected(id: &str) -> (MqttClient, WsPeer, Option) { + let (l, addr) = ws_listener().await; + let client = MqttClient::with_options(ws_opts(id)); + let c = client.clone(); + let connect = tokio::spawn(async move { c.connect(&format!("ws://{addr}/mqtt")).await }); + let (mut peer, offered) = ws_accept(&l).await; + match peer.read(W).await { + WsRead::Packet(p) => assert_eq!(p.kind, CONNECT), + _ => panic!("setup: expected CONNECT over WebSocket"), + } + peer.send_binary(connack(false, 0, &[])).await; + connect.await.expect("join").expect("ws connect"); + (client, peer, offered) + } + + async fn ws_subscribed(id: &str) -> (MqttClient, WsPeer, Arc>>>) { + let (client, mut peer, _) = ws_connected(id).await; + let got = Arc::new(Mutex::new(Vec::new())); + let got2 = Arc::clone(&got); + let c = client.clone(); + let sub = tokio::spawn(async move { + c.subscribe("ws/#", move |m| got2.lock().expect("lock").push(m.payload)) + .await + }); + loop { + match peer.read(W).await { + WsRead::Packet(p) if p.kind == SUBSCRIBE => { + let (pid, f) = parse_filters(&p, true).expect("subscribe"); + peer.send_binary(suback(pid, f.len())).await; + break; + } + WsRead::Packet(_) => {} + _ => panic!("setup: expected SUBSCRIBE over WebSocket"), + } + } + sub.await.expect("join").expect("subscribe"); + (client, peer, got) + } + + async fn delivered(got: &Arc>>>, want: usize) -> Vec> { + let deadline = Instant::now() + Duration::from_secs(2); + while Instant::now() < deadline && got.lock().expect("lock").len() < want { + tokio::time::sleep(Duration::from_millis(20)).await; + } + got.lock().expect("lock").clone() + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn mqtt_6_0_0_3_client_offers_mqtt_subprotocol() { + let (_client, _peer, offered) = ws_connected("d-6003").await; + let offered = offered.unwrap_or_default(); + assert!( + offered.split(',').any(|p| p.trim() == "mqtt"), + "MQTT-6.0.0-3 VIOLATION: offered Sec-WebSocket-Protocol={offered:?}" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn mqtt_6_0_0_1_client_sends_only_binary_frames() { + let (client, mut peer, _) = ws_connected("d-6001a").await; + let c = client.clone(); + let pubs = tokio::spawn(async move { + let _ = c.publish("ws/a", b"x".to_vec()).await; + let _ = c.disconnect().await; + }); + let (_closed, _) = peer.closed_within(W).await; + let _ = pubs.await; + assert!( + peer.non_binary_from_client.is_empty(), + "MQTT-6.0.0-1 VIOLATION: client sent non-binary data frames: {:?}", + peer.non_binary_from_client + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn mqtt_6_0_0_1_text_frame_after_connect_closes_connection() { + let (client, mut peer, _) = ws_connected("d-6001b").await; + let _ = peer.ws.send(Message::Text("not mqtt".into())).await; + let (closed, after) = peer.closed_within(Duration::from_secs(4)).await; + let connected = client.is_connected().await; + assert!( + closed, + "MQTT-6.0.0-1 VIOLATION: after a WebSocket text data frame the client did not close the Network Connection within 4s (is_connected={connected}); MQTT packets sent afterwards={after:?}" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn mqtt_6_0_0_1_text_frame_before_connack_closes_connection() { + let (l, addr) = ws_listener().await; + let client = MqttClient::with_options(ws_opts("d-6001c")); + let c = client.clone(); + let connect = tokio::spawn(async move { c.connect(&format!("ws://{addr}/mqtt")).await }); + let (mut peer, _) = ws_accept(&l).await; + let _ = peer.read(W).await; + let _ = peer.ws.send(Message::Text("not mqtt".into())).await; + tokio::time::sleep(Duration::from_millis(200)).await; + peer.send_binary(connack(false, 0, &[])).await; + let result = timeout(W, connect).await; + let (closed, _) = peer.closed_within(Duration::from_millis(500)).await; + assert!( + closed && !matches!(result, Ok(Ok(Ok(())))), + "MQTT-6.0.0-1 VIOLATION: text data frame before CONNACK was ignored; connect result={result:?}, connection closed={closed}" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn mqtt_6_0_0_2_multiple_packets_in_one_frame() { + let (_client, mut peer, got) = ws_subscribed("d-6002a").await; + let mut both = publish("ws/a", 0, 0, b"one"); + both.extend(publish("ws/b", 0, 0, b"two")); + peer.send_binary(both).await; + let msgs = delivered(&got, 2).await; + assert_eq!( + msgs, + vec![b"one".to_vec(), b"two".to_vec()], + "MQTT-6.0.0-2 VIOLATION: two PUBLISH packets in one WebSocket frame; delivered={msgs:?}" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn mqtt_6_0_0_2_packet_split_across_frames() { + let (client, mut peer, got) = ws_subscribed("d-6002b").await; + let bytes = publish("ws/split", 0, 0, b"payload-split"); + let (a, b) = bytes.split_at(5); + peer.send_binary(a.to_vec()).await; + tokio::time::sleep(Duration::from_millis(50)).await; + peer.send_binary(b.to_vec()).await; + let msgs = delivered(&got, 1).await; + let connected = client.is_connected().await; + assert_eq!( + msgs, + vec![b"payload-split".to_vec()], + "MQTT-6.0.0-2 VIOLATION: PUBLISH split across two WebSocket frames not reassembled; delivered={msgs:?} still_connected={connected}" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn mqtt_6_0_0_2_connack_split_across_frames() { + let (l, addr) = ws_listener().await; + let client = MqttClient::with_options(ws_opts("d-6002c")); + let c = client.clone(); + let connect = tokio::spawn(async move { c.connect(&format!("ws://{addr}/mqtt")).await }); + let (mut peer, _) = ws_accept(&l).await; + let _ = peer.read(W).await; + let bytes = connack(false, 0, &[]); + peer.send_binary(bytes[..2].to_vec()).await; + peer.send_binary(bytes[2..].to_vec()).await; + let result = timeout(W, connect).await; + assert!( + matches!(result, Ok(Ok(Ok(())))), + "MQTT-6.0.0-2: CONNACK split across frames must be reassembled: {result:?}" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn rfc6455_ping_control_frame_does_not_drop_mqtt_session() { + let (client, mut peer, got) = ws_subscribed("d-6000ping").await; + let _ = peer.ws.send(Message::Ping(b"hi".to_vec().into())).await; + tokio::time::sleep(Duration::from_millis(200)).await; + peer.send_binary(publish("ws/after-ping", 0, 0, b"after")) + .await; + let msgs = delivered(&got, 1).await; + let connected = client.is_connected().await; + assert!( + connected && msgs == vec![b"after".to_vec()], + "WebSocket robustness (not an MQTT-6.0.0-1 data frame): a Ping control frame tore down the MQTT session; is_connected={connected} delivered={msgs:?}" + ); + } +} diff --git a/crates/mqtt5/tests/deferred_ack_matrix.rs b/crates/mqtt5/tests/deferred_ack_matrix.rs index dc4f9722..f475f9b1 100644 --- a/crates/mqtt5/tests/deferred_ack_matrix.rs +++ b/crates/mqtt5/tests/deferred_ack_matrix.rs @@ -1,5 +1,4 @@ #![cfg(feature = "broker")] -#![allow(clippy::large_futures)] mod common; @@ -271,10 +270,12 @@ async fn crash_regime_fresh_client_resumes_and_redelivers() { tokio::time::sleep(Duration::from_millis(200)).await; let second = MqttClient::new(&id); - let result = second - .connect_with_options(broker.address(), deferred_options(&id, 8)) - .await - .unwrap(); + let result = Box::pin(second.connect_with_options( + broker.address(), + deferred_options(&id, 8).with_resume_existing_session(true), + )) + .await + .unwrap(); assert!( result.session_present, "the still-running broker resumes the persistent session" @@ -363,8 +364,7 @@ async fn websocket_deferred_qos2_delivers_backpressures_and_acks() { let sub_opts = deferred_options(&client_id("ws-sub"), 1); let subscriber = MqttClient::with_options(sub_opts.clone()); - subscriber - .connect_with_options(broker.address(), sub_opts) + Box::pin(subscriber.connect_with_options(broker.address(), sub_opts)) .await .unwrap(); diff --git a/crates/mqtt5/tests/integration_complete_flow.rs b/crates/mqtt5/tests/integration_complete_flow.rs index 00bde264..9458c53c 100644 --- a/crates/mqtt5/tests/integration_complete_flow.rs +++ b/crates/mqtt5/tests/integration_complete_flow.rs @@ -1,6 +1,4 @@ #![cfg(feature = "broker")] -#![allow(clippy::implicit_clone)] -#![allow(clippy::large_futures)] mod common; @@ -53,14 +51,11 @@ use tokio::sync::Mutex; #[tokio::test] async fn test_complete_mqtt_flow() { - // Start test broker let broker = TestBroker::start().await; - // Create and connect client let client = create_test_client_with_broker("complete-flow", broker.address()).await; assert!(client.is_connected().await); - // Test single subscription and publish using EventCounter let counter = EventCounter::new(); let sub_opts = SubscribeOptions { @@ -73,7 +68,6 @@ async fn test_complete_mqtt_flow() { .await .expect("Failed to subscribe"); - // Publish a message let result = client .publish_qos1("test/topic", b"Hello MQTT") .await @@ -84,36 +78,31 @@ async fn test_complete_mqtt_flow() { PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), } - // Wait for message to be received assert!( counter.wait_for(1, Duration::from_secs(1)).await, "Timeout waiting for message" ); assert_eq!(counter.get(), 1); - // Test unsubscribe client .unsubscribe("test/topic") .await .expect("Failed to unsubscribe"); - // Publish again - should not be received client .publish("test/topic", b"Should not receive") .await .expect("Failed to publish"); tokio::time::sleep(Duration::from_millis(100)).await; - assert_eq!(counter.get(), 1); // Still 1 + assert_eq!(counter.get(), 1); - // Disconnect cleanly client.disconnect().await.expect("Failed to disconnect"); assert!(!client.is_connected().await); } #[tokio::test] async fn test_multiple_subscriptions_and_wildcards() { - // Start test broker let broker = TestBroker::start().await; let client = MqttClient::new(test_client_id("multi-sub")); @@ -123,18 +112,13 @@ async fn test_multiple_subscriptions_and_wildcards() { .await .expect("Failed to connect"); - // Track received messages by topic let messages = Arc::new(Mutex::new(HashMap::>>::new())); - // Subscribe to non-overlapping topics to test wildcard functionality - // without overlapping subscription complications - - // Subscribe to specific topic that won't overlap with wildcards let messages_clone = Arc::clone(&messages); client .subscribe("sensors/exact/temperature", move |msg| { let messages_clone = messages_clone.clone(); - let topic = msg.topic.to_string(); + let topic = msg.topic.clone(); let payload = msg.payload.clone(); tokio::spawn(async move { let mut msgs = messages_clone.lock().await; @@ -144,7 +128,6 @@ async fn test_multiple_subscriptions_and_wildcards() { .await .expect("Failed to subscribe to specific topic"); - // Subscribe to single-level wildcard for different path let messages_clone = Arc::clone(&messages); client .subscribe("devices/+/status", move |msg| { @@ -159,7 +142,6 @@ async fn test_multiple_subscriptions_and_wildcards() { .await .expect("Failed to subscribe to single-level wildcard"); - // Subscribe to multi-level wildcard for different path let messages_clone = Arc::clone(&messages); client .subscribe("system/#", move |msg| { @@ -174,7 +156,6 @@ async fn test_multiple_subscriptions_and_wildcards() { .await .expect("Failed to subscribe to multi-level wildcard"); - // Publish to various topics that match different subscriptions client .publish("sensors/exact/temperature", b"25.5") .await @@ -190,17 +171,13 @@ async fn test_multiple_subscriptions_and_wildcards() { .await .expect("Failed to publish system log"); - // Wait for messages tokio::time::sleep(Duration::from_millis(200)).await; - // Verify received messages let msgs = messages.lock().await; - // Exact subscription should receive only its specific topic assert_eq!(msgs.get("sensors/exact/temperature").unwrap().len(), 1); assert_eq!(msgs.get("sensors/exact/temperature").unwrap()[0], b"25.5"); - // Single-level wildcard should receive device status assert_eq!( msgs.get("wildcard-single:devices/sensor1/status") .unwrap() @@ -212,7 +189,6 @@ async fn test_multiple_subscriptions_and_wildcards() { b"online" ); - // Multi-level wildcard should receive system message assert_eq!( msgs.get("wildcard-multi:system/log/debug").unwrap().len(), 1 @@ -222,7 +198,6 @@ async fn test_multiple_subscriptions_and_wildcards() { b"test message" ); - // Verify no cross-contamination between subscriptions assert!(msgs .get("wildcard-single:sensors/exact/temperature") .is_none()); @@ -236,7 +211,6 @@ async fn test_multiple_subscriptions_and_wildcards() { #[tokio::test] async fn test_qos_levels_and_acknowledgments() { - // Start test broker let broker = TestBroker::start().await; let client = MqttClient::new(test_client_id("qos-test")); @@ -246,14 +220,12 @@ async fn test_qos_levels_and_acknowledgments() { .await .expect("Failed to connect"); - // Test QoS 0 - no packet ID let result = client .publish("test/qos0", b"QoS 0 message") .await .expect("Failed to publish QoS 0"); assert!(matches!(result, PublishResult::QoS0)); - // Test QoS 1 - should get packet ID let result = client .publish_qos1("test/qos1", b"QoS 1 message") .await @@ -263,7 +235,6 @@ async fn test_qos_levels_and_acknowledgments() { PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), } - // Test QoS 2 - should get packet ID let result = client .publish_qos2("test/qos2", b"QoS 2 message") .await @@ -273,11 +244,9 @@ async fn test_qos_levels_and_acknowledgments() { PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), } - // Subscribe and verify QoS downgrade let received_qos = Arc::new(Mutex::new(Vec::new())); let received_qos_clone = Arc::clone(&received_qos); - // Subscribe with QoS 1 let sub_opts = SubscribeOptions { qos: QoS::AtLeastOnce, ..Default::default() @@ -286,7 +255,7 @@ async fn test_qos_levels_and_acknowledgments() { client .subscribe_with_options("qostest/+", sub_opts, move |msg| { let received_qos_clone = received_qos_clone.clone(); - let topic = msg.topic.to_string(); + let topic = msg.topic.clone(); let qos = msg.qos; tokio::spawn(async move { received_qos_clone.lock().await.push((topic, qos)); @@ -295,7 +264,6 @@ async fn test_qos_levels_and_acknowledgments() { .await .expect("Failed to subscribe"); - // Publish with different QoS levels client.publish("qostest/downgrade0", b"msg").await.unwrap(); client .publish_qos1("qostest/downgrade1", b"msg") @@ -311,15 +279,12 @@ async fn test_qos_levels_and_acknowledgments() { let qos_list = received_qos.lock().await; assert_eq!(qos_list.len(), 3); - // QoS 0 stays 0 assert!(qos_list .iter() .any(|(t, q)| t == "qostest/downgrade0" && *q == QoS::AtMostOnce)); - // QoS 1 stays 1 assert!(qos_list .iter() .any(|(t, q)| t == "qostest/downgrade1" && *q == QoS::AtLeastOnce)); - // QoS 2 downgrades to 1 (subscription max) assert!(qos_list .iter() .any(|(t, q)| t == "qostest/downgrade2" && *q == QoS::AtLeastOnce)); @@ -329,19 +294,18 @@ async fn test_qos_levels_and_acknowledgments() { #[tokio::test] async fn test_session_persistence() { - // Start test broker let broker = TestBroker::start().await; let client_id = test_client_id("session-test"); - // First connection with clean_start = false let client1 = MqttClient::new(client_id.clone()); - let mut opts = ConnectOptions::new(client_id.clone()).with_clean_start(false); - opts.properties.session_expiry_interval = Some(300); // 5 minutes + let mut opts = ConnectOptions::new(client_id.clone()) + .with_clean_start(false) + .with_resume_existing_session(true); + opts.properties.session_expiry_interval = Some(300); - let connect_result1 = client1 - .connect_with_options(broker.address(), opts.clone()) + let connect_result1 = Box::pin(client1.connect_with_options(broker.address(), opts.clone())) .await .expect("Failed to connect"); println!( @@ -349,16 +313,13 @@ async fn test_session_persistence() { connect_result1.session_present ); - // Subscribe to a topic client1 .subscribe("persistent/topic", |_| {}) .await .expect("Failed to subscribe"); - // Disconnect client1.disconnect().await.expect("Failed to disconnect"); - // Publish a message while disconnected (from another client) let publisher = MqttClient::new(test_client_id("publisher")); publisher @@ -374,11 +335,9 @@ async fn test_session_persistence() { .await .expect("Publisher failed to disconnect"); - // Reconnect with same client ID and clean_start = false let client2 = MqttClient::new(client_id); - let connect_result2 = client2 - .connect_with_options(broker.address(), opts) + let connect_result2 = Box::pin(client2.connect_with_options(broker.address(), opts)) .await .expect("Failed to reconnect"); println!( @@ -389,18 +348,11 @@ async fn test_session_persistence() { let received = Arc::new(AtomicU32::new(0)); let received_clone = Arc::clone(&received); - // DON'T re-subscribe! The subscription should be restored from the persistent session - // Just set up the callback on the existing subscription (if the client supports this) - // For now, let's wait for any messages that might be delivered from the restored session - - // Wait longer to see if the offline message is delivered automatically println!("Waiting for offline message from restored session..."); tokio::time::sleep(Duration::from_secs(1)).await; let count = received.load(Ordering::SeqCst); println!("Received message count after waiting: {count}"); - // Since we can't set up callbacks without subscribing in our current architecture, - // let's try re-subscribing with the callback client2 .subscribe("persistent/topic", move |msg| { println!( @@ -413,21 +365,16 @@ async fn test_session_persistence() { .await .expect("Failed to re-subscribe"); - // Wait again after re-subscribing println!("Waiting for message after re-subscribe..."); tokio::time::sleep(Duration::from_millis(500)).await; let final_count = received.load(Ordering::SeqCst); println!("Final received message count: {final_count}"); - // NOTE: Session persistence behavior varies by broker implementation - // Some brokers may not queue QoS 1 messages for offline persistent sessions - // This test verifies that session_present=true works, which is the key requirement if final_count == 0 { println!( "BROKER NOTE: This broker doesn't queue QoS 1 messages for offline persistent sessions" ); println!("The session persistence mechanism itself works (session_present=true)"); - // Test passes - session persistence is working, message queueing is broker-dependent } else { assert_eq!(final_count, 1); } @@ -437,7 +384,6 @@ async fn test_session_persistence() { #[tokio::test] async fn test_publish_options_and_properties() { - // Start test broker let broker = TestBroker::start().await; let client = MqttClient::new(test_client_id("pub-options")); @@ -460,7 +406,6 @@ async fn test_publish_options_and_properties() { .await .expect("Failed to subscribe"); - // Publish with various options let user_properties = vec![ ("key1".to_string(), "value1".to_string()), ("key2".to_string(), "value2".to_string()), @@ -487,21 +432,15 @@ async fn test_publish_options_and_properties() { .await .expect("Failed to publish with options"); - // Wait and verify tokio::time::sleep(Duration::from_millis(200)).await; let msgs = messages.lock().await; - // We might receive both normal delivery and retained delivery - // Take the first message (normal delivery) for properties testing assert!(!msgs.is_empty()); let msg = &msgs[0]; assert_eq!(msg.topic, "test/properties"); assert_eq!(&msg.payload[..], b"Message with properties"); - // Note: The first message is the normal delivery (retain: false) - // The retained message will be tested with a new subscription below - // Test retained message delivery to new subscriber let retained_received = Arc::new(AtomicU32::new(0)); let retained_clone = Arc::clone(&retained_received); @@ -514,11 +453,9 @@ async fn test_publish_options_and_properties() { .await .expect("Failed to resubscribe"); - // Should immediately receive the retained message tokio::time::sleep(Duration::from_millis(100)).await; assert_eq!(retained_received.load(Ordering::SeqCst), 1); - // Clear retained message client .publish("test/properties", b"") .await @@ -529,7 +466,6 @@ async fn test_publish_options_and_properties() { #[tokio::test] async fn test_subscription_options() { - // Start test broker let broker = TestBroker::start().await; let client = MqttClient::new(test_client_id("sub-options")); @@ -539,21 +475,13 @@ async fn test_subscription_options() { .await .expect("Failed to connect"); - // Skip No Local option test - not yet implemented in broker - // The No Local option prevents delivery of messages published by the same client. - // This is an MQTT v5 feature that needs to be implemented in the broker. - - // Test Retain Handling options - // First, set a retained message client .publish_retain("test/retain/handling", b"Retained message") .await .expect("Failed to publish retained"); - // Give broker time to process retained message tokio::time::sleep(Duration::from_millis(100)).await; - // Subscribe with SEND_AT_SUBSCRIBE (default) let received_retained = Arc::new(AtomicU32::new(0)); let received_clone = Arc::clone(&received_retained); @@ -572,7 +500,6 @@ async fn test_subscription_options() { println!("Received {count} retained messages"); assert_eq!(count, 1); - // Clear retained message client.publish("test/retain/handling", b"").await.unwrap(); client.disconnect().await.expect("Failed to disconnect"); @@ -580,7 +507,6 @@ async fn test_subscription_options() { #[tokio::test] async fn test_large_payload_handling() { - // Start test broker let broker = TestBroker::start().await; let client = MqttClient::new(test_client_id("large-payload")); @@ -590,7 +516,6 @@ async fn test_large_payload_handling() { .await .expect("Failed to connect"); - // Create a large payload (1MB) let large_payload = vec![0x42; 1024 * 1024]; let payload_clone = large_payload.clone(); @@ -608,13 +533,11 @@ async fn test_large_payload_handling() { .await .expect("Failed to subscribe"); - // Publish large message client .publish_qos1("test/large", large_payload.clone()) .await .expect("Failed to publish large message"); - // Wait and verify tokio::time::sleep(Duration::from_millis(500)).await; let received_payload = received.lock().await; @@ -626,7 +549,6 @@ async fn test_large_payload_handling() { #[tokio::test] async fn test_concurrent_operations() { - // Start test broker let broker = TestBroker::start().await; let broker_addr = broker.address().to_string(); @@ -639,7 +561,6 @@ async fn test_concurrent_operations() { let received = Arc::new(AtomicU32::new(0)); - // Subscribe to multiple topics for i in 0..10 { let received_clone = Arc::clone(&received); client @@ -650,7 +571,6 @@ async fn test_concurrent_operations() { .expect("Failed to subscribe"); } - // Spawn multiple publishers let mut handles = vec![]; for i in 0..10 { @@ -669,15 +589,12 @@ async fn test_concurrent_operations() { handles.push(handle); } - // Wait for all publishers for handle in handles { handle.await.expect("Publisher task failed"); } - // Wait for all messages tokio::time::sleep(Duration::from_millis(500)).await; - // Should have received 100 messages (10 topics × 10 messages) assert_eq!(received.load(Ordering::SeqCst), 100); client.disconnect().await.expect("Failed to disconnect"); diff --git a/crates/mqtt5/tests/persistence.rs b/crates/mqtt5/tests/persistence.rs index 36648aed..8fd1adb2 100644 --- a/crates/mqtt5/tests/persistence.rs +++ b/crates/mqtt5/tests/persistence.rs @@ -1,5 +1,4 @@ #![cfg(feature = "broker")] -#![allow(clippy::large_futures)] mod common; use common::TestBroker; @@ -12,40 +11,34 @@ use tokio::time::sleep; #[tokio::test] async fn test_clean_start_true() { - // Start test broker let broker = TestBroker::start().await; let options = ConnectOptions::new("clean-start-true").with_clean_start(true); let client = MqttClient::with_options(options); - // First connection - let session_present = client - .connect_with_options( - broker.address(), - ConnectOptions::new("clean-start-true").with_clean_start(true), - ) - .await - .unwrap(); + let session_present = Box::pin(client.connect_with_options( + broker.address(), + ConnectOptions::new("clean-start-true").with_clean_start(true), + )) + .await + .unwrap(); assert!( !session_present.session_present, "First connection should not have session present" ); - // Subscribe to a topic client.subscribe("test/clean", |_| {}).await.unwrap(); client.disconnect().await.unwrap(); - // Second connection with clean_start=true - let session_present = client - .connect_with_options( - broker.address(), - ConnectOptions::new("clean-start-true").with_clean_start(true), - ) - .await - .unwrap(); + let session_present = Box::pin(client.connect_with_options( + broker.address(), + ConnectOptions::new("clean-start-true").with_clean_start(true), + )) + .await + .unwrap(); assert!( !session_present.session_present, @@ -57,45 +50,44 @@ async fn test_clean_start_true() { #[tokio::test] async fn test_clean_start_false() { - // Start test broker let broker = TestBroker::start().await; let client_id = "persist-test-1"; - // First connection with clean_start=true to ensure clean slate let client1 = MqttClient::with_options(ConnectOptions::new(client_id).with_clean_start(true)); client1.connect(broker.address()).await.unwrap(); - // Subscribe to topics client1.subscribe("test/persist/1", |_| {}).await.unwrap(); client1.subscribe("test/persist/2", |_| {}).await.unwrap(); client1.disconnect().await.unwrap(); - // Second connection with clean_start=false - let client2 = MqttClient::with_options(ConnectOptions::new(client_id).with_clean_start(false)); + let client2 = MqttClient::with_options( + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); - let session_present = client2 - .connect_with_options( + let session_present = Box::pin( + client2.connect_with_options( broker.address(), - ConnectOptions::new(client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); - // Note: Some brokers may not preserve sessions even with clean_start=false let session_present_flag = session_present.session_present; println!("Session present: {session_present_flag}"); if !session_present.session_present { println!("Warning: Broker did not preserve session. This is broker-dependent behavior."); } - // Subscriptions should still be active - // Test by publishing to the subscribed topics let received = Arc::new(AtomicU32::new(0)); let received_clone = received.clone(); - // Re-subscribe to set up callback (broker maintains subscription but we need local callback) client2 .subscribe("test/persist/1", move |_| { received_clone.fetch_add(1, Ordering::Relaxed); @@ -106,7 +98,6 @@ async fn test_clean_start_false() { client2.publish("test/persist/1", "test").await.unwrap(); sleep(Duration::from_millis(500)).await; - // Only check if session was actually preserved if session_present.session_present { assert!( received.load(Ordering::Relaxed) > 0, @@ -119,31 +110,26 @@ async fn test_clean_start_false() { #[tokio::test] async fn test_session_expiry_interval() { - // Start test broker let broker = TestBroker::start().await; let client_id = "session-expiry-test"; - // Connect with session expiry interval let options = ConnectOptions::new(client_id) .with_clean_start(false) - .with_session_expiry_interval(5); // 5 seconds + .with_resume_existing_session(true) + .with_session_expiry_interval(5); let client1 = MqttClient::with_options(options.clone()); client1.connect(broker.address()).await.unwrap(); - // Subscribe to a topic client1.subscribe("test/expiry", |_| {}).await.unwrap(); client1.disconnect().await.unwrap(); - // Wait less than expiry interval sleep(Duration::from_secs(2)).await; - // Reconnect - session should still exist let client2 = MqttClient::with_options(options.clone()); - let session_present = client2 - .connect_with_options(broker.address(), options.clone()) + let session_present = Box::pin(client2.connect_with_options(broker.address(), options.clone())) .await .unwrap(); @@ -153,20 +139,20 @@ async fn test_session_expiry_interval() { ); client2.disconnect().await.unwrap(); - // Wait for session to expire sleep(Duration::from_secs(4)).await; - // Reconnect - session should be gone let client3 = MqttClient::with_options(options); - let session_present = client3 - .connect_with_options( + let session_present = Box::pin( + client3.connect_with_options( broker.address(), - ConnectOptions::new(client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); - // Broker might not have expired it yet, so we don't assert here println!( "Session present after expiry: {}", session_present.session_present @@ -177,14 +163,14 @@ async fn test_session_expiry_interval() { #[tokio::test] async fn test_qos1_message_persistence() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("persist-pub"); let sub_client_id = "persist-sub-qos1"; - // Subscriber connects and subscribes - let sub_options = ConnectOptions::new(sub_client_id).with_clean_start(false); + let sub_options = ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true); let sub_client = MqttClient::with_options(sub_options); sub_client.connect(broker.address()).await.unwrap(); @@ -200,10 +186,8 @@ async fn test_qos1_message_persistence() { .await .unwrap(); - // Disconnect subscriber sub_client.disconnect().await.unwrap(); - // Publisher sends QoS 1 messages while subscriber is offline pub_client.connect(broker.address()).await.unwrap(); for i in 0..5 { @@ -215,20 +199,25 @@ async fn test_qos1_message_persistence() { pub_client.disconnect().await.unwrap(); - // Subscriber reconnects let received = Arc::new(AtomicU32::new(0)); let received_clone = received.clone(); - let sub_client2 = - MqttClient::with_options(ConnectOptions::new(sub_client_id).with_clean_start(false)); + let sub_client2 = MqttClient::with_options( + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); - let session_present = sub_client2 - .connect_with_options( + let session_present = Box::pin( + sub_client2.connect_with_options( broker.address(), - ConnectOptions::new(sub_client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); println!( "Session present after reconnect: {}", @@ -238,7 +227,6 @@ async fn test_qos1_message_persistence() { println!("Warning: Broker did not restore session for QoS persistence test"); } - // Re-subscribe to set callback sub_client2 .subscribe_with_options( "test/persist/qos1", @@ -253,12 +241,10 @@ async fn test_qos1_message_persistence() { .await .unwrap(); - // Wait for queued messages sleep(Duration::from_secs(2)).await; let count = received.load(Ordering::Relaxed); println!("Received {count} offline messages"); - // Only assert if session was preserved if session_present.session_present { assert!( count > 0, @@ -274,14 +260,14 @@ async fn test_qos1_message_persistence() { #[tokio::test] async fn test_qos2_message_persistence() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("persist-pub-qos2"); let sub_client_id = "persist-sub-qos2"; - // Subscriber connects and subscribes with QoS 2 - let sub_options = ConnectOptions::new(sub_client_id).with_clean_start(false); + let sub_options = ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true); let sub_client = MqttClient::with_options(sub_options); sub_client.connect(broker.address()).await.unwrap(); @@ -297,10 +283,8 @@ async fn test_qos2_message_persistence() { .await .unwrap(); - // Disconnect subscriber sub_client.disconnect().await.unwrap(); - // Publisher sends QoS 2 messages while subscriber is offline pub_client.connect(broker.address()).await.unwrap(); for i in 0..3 { @@ -312,16 +296,17 @@ async fn test_qos2_message_persistence() { pub_client.disconnect().await.unwrap(); - // Subscriber reconnects let messages = Arc::new(std::sync::Mutex::new(Vec::new())); let messages_clone = messages.clone(); - let sub_client2 = - MqttClient::with_options(ConnectOptions::new(sub_client_id).with_clean_start(false)); + let sub_client2 = MqttClient::with_options( + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); sub_client2.connect(broker.address()).await.unwrap(); - // Re-subscribe to set callback sub_client2 .subscribe_with_options( "test/persist/qos2", @@ -339,7 +324,6 @@ async fn test_qos2_message_persistence() { .await .unwrap(); - // Wait for queued messages sleep(Duration::from_secs(2)).await; { @@ -347,12 +331,11 @@ async fn test_qos2_message_persistence() { let msg_count = msgs.len(); println!("Received {msg_count} QoS 2 offline messages"); - // Should receive exactly once let mut unique_msgs = msgs.clone(); unique_msgs.sort(); unique_msgs.dedup(); assert_eq!(msgs.len(), unique_msgs.len(), "No duplicate QoS 2 messages"); - } // Drop the lock before awaiting + } match sub_client2.disconnect().await { Ok(()) | Err(mqtt5::MqttError::NotConnected) => {} @@ -362,12 +345,10 @@ async fn test_qos2_message_persistence() { #[tokio::test] async fn test_subscription_persistence() { - // Start test broker let broker = TestBroker::start().await; let client_id = "sub-persist-test"; - // First connection - subscribe to multiple topics let client1 = MqttClient::with_options(ConnectOptions::new(client_id).with_clean_start(true)); client1.connect(broker.address()).await.unwrap(); @@ -377,19 +358,25 @@ async fn test_subscription_persistence() { client1.disconnect().await.unwrap(); - // Second connection - subscriptions should persist let received_topics = Arc::new(std::sync::Mutex::new(Vec::new())); let received_topics_clone = received_topics.clone(); - let client2 = MqttClient::with_options(ConnectOptions::new(client_id).with_clean_start(false)); + let client2 = MqttClient::with_options( + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); - let session_present = client2 - .connect_with_options( + let session_present = Box::pin( + client2.connect_with_options( broker.address(), - ConnectOptions::new(client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); println!( "Session present for subscription persistence: {}", @@ -397,13 +384,10 @@ async fn test_subscription_persistence() { ); if !session_present.session_present { println!("Warning: Broker did not preserve session for subscription test"); - // Skip the rest of the test if session wasn't preserved client2.disconnect().await.unwrap(); return; } - // Need to re-subscribe to set local callbacks - // (broker maintains subscriptions but we need local handlers) client2 .subscribe("test/sub/+", move |msg| { received_topics_clone @@ -414,7 +398,6 @@ async fn test_subscription_persistence() { .await .unwrap(); - // Publish to subscribed topics client2.publish("test/sub/1", "msg1").await.unwrap(); client2.publish("test/sub/2", "msg2").await.unwrap(); client2.publish("test/sub/3", "msg3").await.unwrap(); @@ -427,20 +410,18 @@ async fn test_subscription_persistence() { topics.len() >= 3, "Should receive messages on persisted subscriptions" ); - } // Drop the lock before awaiting + } client2.disconnect().await.unwrap(); } #[tokio::test] async fn test_will_message_persistence() { - // Start test broker let broker = TestBroker::start().await; let will_client_id = "will-persist-test"; let sub_client = MqttClient::new("will-sub"); - // Subscribe to will topic sub_client.connect(broker.address()).await.unwrap(); let will_received = Arc::new(AtomicBool::new(false)); @@ -457,26 +438,22 @@ async fn test_will_message_persistence() { .await .unwrap(); - // Connect with will message and persistent session let will_msg = mqtt5::WillMessage::new("test/will/persist", "Client died") .with_qos(QoS::AtLeastOnce) .with_retain(false); let will_options = ConnectOptions::new(will_client_id) .with_clean_start(false) + .with_resume_existing_session(true) .with_will(will_msg); let will_client = MqttClient::with_options(will_options); will_client.connect(broker.address()).await.unwrap(); - // Simulate abnormal disconnection by dropping the client - // This causes the TCP connection to close without sending DISCONNECT drop(will_client); - // Wait for will message sleep(Duration::from_secs(2)).await; - // Will message delivery depends on broker implementation let received = will_received.load(Ordering::Relaxed); println!("Will message received: {received}"); if !received { @@ -488,18 +465,17 @@ async fn test_will_message_persistence() { #[tokio::test] async fn test_packet_id_persistence() { - // Start test broker let broker = TestBroker::start().await; - // Test that packet IDs are managed correctly across reconnections let client_id = "packet-id-persist"; - let options = ConnectOptions::new(client_id).with_clean_start(false); + let options = ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true); let client1 = MqttClient::with_options(options.clone()); client1.connect(broker.address()).await.unwrap(); - // Send some QoS 1 messages to allocate packet IDs let mut first_ids = Vec::new(); for i in 0..5 { let id = client1 @@ -511,7 +487,6 @@ async fn test_packet_id_persistence() { client1.disconnect().await.unwrap(); - // Reconnect and send more messages let client2 = MqttClient::with_options(options); client2.connect(broker.address()).await.unwrap(); @@ -524,8 +499,6 @@ async fn test_packet_id_persistence() { second_ids.push(id); } - // With clean disconnection, packet IDs can be reused - // This is normal behavior as the previous IDs were acknowledged println!("First session IDs: {first_ids:?}"); println!("Second session IDs: {second_ids:?}"); @@ -534,16 +507,16 @@ async fn test_packet_id_persistence() { #[tokio::test] async fn test_inflight_message_persistence() { - // Start test broker let broker = TestBroker::start().await; - // Test that in-flight QoS 1/2 messages are retransmitted after reconnection let pub_client_id = "inflight-pub"; let sub_client_id = "inflight-sub"; - // Set up subscriber - let sub_client = - MqttClient::with_options(ConnectOptions::new(sub_client_id).with_clean_start(false)); + let sub_client = MqttClient::with_options( + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); sub_client.connect(broker.address()).await.unwrap(); let received = Arc::new(AtomicU32::new(0)); @@ -563,40 +536,38 @@ async fn test_inflight_message_persistence() { .await .unwrap(); - // Publisher sends messages - let pub_client = - MqttClient::with_options(ConnectOptions::new(pub_client_id).with_clean_start(false)); + let pub_client = MqttClient::with_options( + ConnectOptions::new(pub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); pub_client.connect(broker.address()).await.unwrap(); - // Send QoS 1 messages rapidly then disconnect - // Some might still be in-flight for i in 0..10 { let _ = pub_client .publish_qos1("test/inflight", format!("Msg {i}")) .await; } - // Quick disconnect might leave some messages in-flight pub_client.disconnect().await.unwrap(); - // Wait a bit sleep(Duration::from_millis(500)).await; let initial_count = received.load(Ordering::Relaxed); println!("Initially received: {initial_count} messages"); - // Reconnect publisher - any in-flight messages should be retransmitted - let pub_client2 = - MqttClient::with_options(ConnectOptions::new(pub_client_id).with_clean_start(false)); + let pub_client2 = MqttClient::with_options( + ConnectOptions::new(pub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); pub_client2.connect(broker.address()).await.unwrap(); - // Wait for potential retransmissions sleep(Duration::from_secs(1)).await; let final_count = received.load(Ordering::Relaxed); println!("Finally received: {final_count} messages"); - // Should eventually receive all messages assert!( final_count >= 10, "Should receive all messages including retransmissions" @@ -613,8 +584,11 @@ async fn test_qos2_outbound_inflight_resend_on_reconnect() { let pub_client = MqttClient::new("qos2-inflight-pub"); let sub_client_id = "qos2-inflight-sub"; - let sub_client = - MqttClient::with_options(ConnectOptions::new(sub_client_id).with_clean_start(false)); + let sub_client = MqttClient::with_options( + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); sub_client.connect(broker.address()).await.unwrap(); sub_client @@ -643,16 +617,22 @@ async fn test_qos2_outbound_inflight_resend_on_reconnect() { let messages = Arc::new(std::sync::Mutex::new(Vec::new())); let messages_clone = messages.clone(); - let sub_client2 = - MqttClient::with_options(ConnectOptions::new(sub_client_id).with_clean_start(false)); + let sub_client2 = MqttClient::with_options( + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); - let session = sub_client2 - .connect_with_options( + let session = Box::pin( + sub_client2.connect_with_options( broker.address(), - ConnectOptions::new(sub_client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); sub_client2 .subscribe_with_options( @@ -704,8 +684,11 @@ async fn test_clean_start_clears_inflight_state() { let pub_client = MqttClient::new("clean-inflight-pub"); let sub_client_id = "clean-inflight-sub"; - let sub_client = - MqttClient::with_options(ConnectOptions::new(sub_client_id).with_clean_start(false)); + let sub_client = MqttClient::with_options( + ConnectOptions::new(sub_client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ); sub_client.connect(broker.address()).await.unwrap(); sub_client @@ -737,13 +720,12 @@ async fn test_clean_start_clears_inflight_state() { let sub_client2 = MqttClient::with_options(ConnectOptions::new(sub_client_id).with_clean_start(true)); - let session = sub_client2 - .connect_with_options( - broker.address(), - ConnectOptions::new(sub_client_id).with_clean_start(true), - ) - .await - .unwrap(); + let session = Box::pin(sub_client2.connect_with_options( + broker.address(), + ConnectOptions::new(sub_client_id).with_clean_start(true), + )) + .await + .unwrap(); assert!( !session.session_present, diff --git a/crates/mqtt5/tests/retained_messages.rs b/crates/mqtt5/tests/retained_messages.rs deleted file mode 100644 index f67a0de3..00000000 --- a/crates/mqtt5/tests/retained_messages.rs +++ /dev/null @@ -1,228 +0,0 @@ -// DISABLED: These tests require direct access to session state which is not -// exposed in the public API. Retained message functionality is tested through -// actual MQTT behavior in: -// - integration_complete_flow.rs (tests retained message delivery to new subscribers) -// - client_publish.rs::test_publish_retain (tests the publish_retain method) -// - message_queuing.rs::test_retained_message_queuing (tests retained message queuing) -// -// To re-enable these tests, they would need to be moved to unit tests in the session module. - -/* -use mqtt5::{MqttClient, MqttError, QoS}; -use std::sync::Arc; -use tokio::sync::Mutex; - -#[tokio::test] -async fn test_retained_message_storage_and_retrieval() { - let client = MqttClient::new("test-retained-client"); - - // Store messages received - let received_messages = Arc::new(Mutex::new(Vec::new())); - let received_clone = Arc::clone(&received_messages); - - // Subscribe to a topic first (to ensure callback is registered) - let subscribe_result = client.subscribe("test/retained/topic", move |msg| { - let received = received_clone.clone(); - tokio::spawn(async move { - let mut msgs = received.lock().await; - msgs.push((msg.topic.clone(), msg.payload.clone(), msg.retain)); - }); - }).await; - - // This will fail since we're not connected, but that's ok for this unit test - assert!(subscribe_result.is_err()); - - // Test retained message storage in session - let session = client.session_state().await; - - // Create a retained message - let packet = mqtt5::packet::publish::PublishPacket { - topic_name: "test/retained/topic".to_string(), - payload: b"retained message".to_vec().into(), - qos: QoS::AtLeastOnce, - retain: true, - dup: false, - packet_id: None, - properties: Default::default(), - }; - - // Store the retained message - session.store_retained_message(&packet).await; - - // Retrieve retained messages for exact topic - let retained = session.get_retained_messages("test/retained/topic").await; - assert_eq!(retained.len(), 1); - assert_eq!(retained[0].topic, "test/retained/topic"); - assert_eq!(&retained[0].payload[..], b"retained message"); - assert_eq!(retained[0].qos, QoS::AtLeastOnce); - - // Clear retained message with empty payload - let clear_packet = mqtt5::packet::publish::PublishPacket { - topic_name: "test/retained/topic".to_string(), - payload: vec![].into(), - qos: QoS::AtMostOnce, - retain: true, - dup: false, - packet_id: None, - properties: Default::default(), - }; - - session.store_retained_message(&clear_packet).await; - - // Verify message was cleared - let retained = session.get_retained_messages("test/retained/topic").await; - assert_eq!(retained.len(), 0); -} - -#[tokio::test] -async fn test_retained_message_wildcard_matching() { - let client = MqttClient::new("test-retained-wildcard"); - let session = client.session_state().await; - - // Store multiple retained messages - let topics = vec![ - ("home/room1/temperature", "20.5".as_bytes()), - ("home/room1/humidity", "65".as_bytes()), - ("home/room2/temperature", "22.0".as_bytes()), - ("home/room2/humidity", "60".as_bytes()), - ("office/room1/temperature", "21.0".as_bytes()), - ]; - - for (topic, payload) in topics { - let packet = mqtt5::packet::publish::PublishPacket { - topic_name: topic.to_string(), - payload: payload.to_vec(), - qos: QoS::AtMostOnce, - retain: true, - dup: false, - packet_id: None, - properties: Default::default(), - }; - session.store_retained_message(&packet).await; - } - - // Test single-level wildcard - let retained = session.get_retained_messages("home/+/temperature").await; - assert_eq!(retained.len(), 2); - assert!(retained.iter().any(|m| m.topic == "home/room1/temperature")); - assert!(retained.iter().any(|m| m.topic == "home/room2/temperature")); - - // Test multi-level wildcard - let retained = session.get_retained_messages("home/#").await; - assert_eq!(retained.len(), 4); - - // Test specific topic - let retained = session.get_retained_messages("office/room1/temperature").await; - assert_eq!(retained.len(), 1); - assert_eq!(retained[0].payload, "21.0".as_bytes()); -} - -#[tokio::test] -async fn test_retained_message_qos_preservation() { - let client = MqttClient::new("test-retained-qos"); - let session = client.session_state().await; - - // Store retained messages with different QoS levels - let qos_levels = vec![ - (QoS::AtMostOnce, "test/qos0"), - (QoS::AtLeastOnce, "test/qos1"), - (QoS::ExactlyOnce, "test/qos2"), - ]; - - for (qos, topic) in qos_levels { - let packet = mqtt5::packet::publish::PublishPacket { - topic_name: topic.to_string(), - payload: format!("QoS {:?} message", qos).into_bytes(), - qos, - retain: true, - dup: false, - packet_id: None, - properties: Default::default(), - }; - session.store_retained_message(&packet).await; - } - - // Verify QoS is preserved - let retained = session.get_retained_messages("test/qos0").await; - assert_eq!(retained[0].qos, QoS::AtMostOnce); - - let retained = session.get_retained_messages("test/qos1").await; - assert_eq!(retained[0].qos, QoS::AtLeastOnce); - - let retained = session.get_retained_messages("test/qos2").await; - assert_eq!(retained[0].qos, QoS::ExactlyOnce); -} - -#[tokio::test] -async fn test_publish_retain_method() { - let client = MqttClient::new("test-publish-retain"); - - // Try to publish retained message (will fail without connection) - let result = client.publish_retain("test/topic", "retained data").await; - assert!(matches!(result, Err(MqttError::NotConnected))); -} - -#[tokio::test] -async fn test_retained_message_store_isolation() { - // Create two separate clients - let client1 = MqttClient::new("client1"); - let client2 = MqttClient::new("client2"); - - let session1 = client1.session_state().await; - let session2 = client2.session_state().await; - - // Store retained message in client1's session - let packet = mqtt5::packet::publish::PublishPacket { - topic_name: "test/isolated".to_string(), - payload: b"client1 message".to_vec().into(), - qos: QoS::AtMostOnce, - retain: true, - dup: false, - packet_id: None, - properties: Default::default(), - }; - session1.store_retained_message(&packet).await; - - // Verify client1 has the message - let retained1 = session1.get_retained_messages("test/isolated").await; - assert_eq!(retained1.len(), 1); - - // Verify client2 does NOT have the message (sessions are isolated) - let retained2 = session2.get_retained_messages("test/isolated").await; - assert_eq!(retained2.len(), 0); -} - -#[tokio::test] -async fn test_retained_message_properties() { - let client = MqttClient::new("test-retained-props"); - let session = client.session_state().await; - - // Create a retained message with properties - let mut properties = mqtt5::protocol::v5::properties::Properties::default(); - let _ = properties.add( - mqtt5::protocol::v5::properties::PropertyId::MessageExpiryInterval, - mqtt5::protocol::v5::properties::PropertyValue::FourByteInteger(3600), - ); - - let packet = mqtt5::packet::publish::PublishPacket { - topic_name: "test/with/properties".to_string(), - payload: b"message with props".to_vec().into(), - qos: QoS::AtLeastOnce, - retain: true, - dup: false, - packet_id: None, - properties, - }; - - session.store_retained_message(&packet).await; - - // Retrieve and verify properties are preserved - let retained = session.get_retained_messages("test/with/properties").await; - assert_eq!(retained.len(), 1); - - let msg_expiry = retained[0].properties.get( - mqtt5::protocol::v5::properties::PropertyId::MessageExpiryInterval - ); - assert!(msg_expiry.is_some()); -} -*/ diff --git a/crates/mqtt5/tests/session_security.rs b/crates/mqtt5/tests/session_security.rs index 0df89d83..d79101ae 100644 --- a/crates/mqtt5/tests/session_security.rs +++ b/crates/mqtt5/tests/session_security.rs @@ -1,5 +1,4 @@ #![cfg(feature = "broker")] -#![allow(clippy::large_futures)] mod common; @@ -94,8 +93,7 @@ async fn test_session_user_binding_rejects_different_user() { .with_credentials("alice", b"pass1") .with_session_expiry_interval(300); let alice = MqttClient::with_options(alice_opts.clone()); - alice - .connect_with_options(broker.address(), alice_opts) + Box::pin(alice.connect_with_options(broker.address(), alice_opts)) .await .expect("alice connect"); alice.subscribe("test/bind", |_| {}).await.unwrap(); @@ -103,10 +101,11 @@ async fn test_session_user_binding_rejects_different_user() { let bob_opts = ConnectOptions::new(shared_client_id) .with_clean_start(false) + .with_resume_existing_session(true) .with_credentials("bob", b"pass2") .with_session_expiry_interval(300); let bob = MqttClient::with_options(bob_opts.clone()); - let result = bob.connect_with_options(broker.address(), bob_opts).await; + let result = Box::pin(bob.connect_with_options(broker.address(), bob_opts)).await; assert!( result.is_err(), @@ -128,8 +127,7 @@ async fn test_session_user_binding_allows_same_user() { .with_credentials("alice", b"pass1") .with_session_expiry_interval(300); let client1 = MqttClient::with_options(opts.clone()); - client1 - .connect_with_options(broker.address(), opts) + Box::pin(client1.connect_with_options(broker.address(), opts)) .await .expect("first connect"); client1.subscribe("test/same", |_| {}).await.unwrap(); @@ -137,11 +135,11 @@ async fn test_session_user_binding_allows_same_user() { let resume_opts = ConnectOptions::new(client_id) .with_clean_start(false) + .with_resume_existing_session(true) .with_credentials("alice", b"pass1") .with_session_expiry_interval(300); let client2 = MqttClient::with_options(resume_opts.clone()); - let result = client2 - .connect_with_options(broker.address(), resume_opts) + let result = Box::pin(client2.connect_with_options(broker.address(), resume_opts)) .await .expect("same user reconnect must succeed"); @@ -217,8 +215,7 @@ async fn test_session_resume_preserves_subscriptions_with_acl() { .with_credentials("alice", b"pass1") .with_session_expiry_interval(300); let client1 = MqttClient::with_options(opts.clone()); - client1 - .connect_with_options(broker.address(), opts) + Box::pin(client1.connect_with_options(broker.address(), opts)) .await .expect("first connect"); @@ -227,11 +224,11 @@ async fn test_session_resume_preserves_subscriptions_with_acl() { let resume_opts = ConnectOptions::new(client_id) .with_clean_start(false) + .with_resume_existing_session(true) .with_credentials("alice", b"pass1") .with_session_expiry_interval(300); let client2 = MqttClient::with_options(resume_opts.clone()); - let result = client2 - .connect_with_options(broker.address(), resume_opts) + let result = Box::pin(client2.connect_with_options(broker.address(), resume_opts)) .await .expect("reconnect must succeed"); @@ -250,8 +247,7 @@ async fn test_session_resume_preserves_subscriptions_with_acl() { .with_clean_start(true) .with_credentials("alice", b"pass1"); let publisher = MqttClient::with_options(pub_opts.clone()); - publisher - .connect_with_options(broker.address(), pub_opts) + Box::pin(publisher.connect_with_options(broker.address(), pub_opts)) .await .expect("publisher connect"); diff --git a/crates/mqtt5/tests/session_state_property_tests.rs b/crates/mqtt5/tests/session_state_property_tests.rs index 7cb0676a..2467c27d 100644 --- a/crates/mqtt5/tests/session_state_property_tests.rs +++ b/crates/mqtt5/tests/session_state_property_tests.rs @@ -8,9 +8,6 @@ //! - Unacked message tracking //! - `QoS` state consistency //! - Subscription management -#![allow(clippy::cast_possible_truncation)] -#![allow(clippy::cast_precision_loss)] -#![allow(deprecated)] use mqtt5::packet::publish::PublishPacket; use mqtt5::packet::subscribe::{RetainHandling, SubscriptionOptions}; @@ -38,12 +35,7 @@ fn qos_level() -> impl Strategy { /// Generate session expiry intervals fn session_expiry() -> impl Strategy { - prop_oneof![ - Just(0), // Immediate expiry - 1..3600u32, // 1 second to 1 hour - Just(86400), // 1 day - Just(u32::MAX), // Maximum value - ] + prop_oneof![Just(0), 1..3600u32, Just(86400), Just(u32::MAX),] } /// Generate publish packets for testing @@ -84,7 +76,6 @@ mod clean_start_tests { let session = SessionState::new("client1".to_string(), config.clone(), false); - // Add some subscriptions let sub = Subscription { topic_filter: "test/+".to_string(), options: SubscriptionOptions { @@ -96,7 +87,6 @@ mod clean_start_tests { }; session.add_subscription(sub.topic_filter.clone(), sub).await.unwrap(); - // Add some unacked messages for &id in &packet_ids { let packet = publish_packet(id, qos); session.store_unacked_publish(packet).await.unwrap(); @@ -108,10 +98,8 @@ mod clean_start_tests { prop_assert!(!initial_subs.is_empty()); prop_assert!(!initial_unacked.is_empty()); - // Create new session with clean_start = true let clean_session = SessionState::new("client1".to_string(), config, true); - // Clean session should have no state let clean_subs = clean_session.all_subscriptions().await; let clean_unacked = clean_session.get_unacked_publishes().await; @@ -133,10 +121,8 @@ mod clean_start_tests { ..SessionConfig::default() }; - // Sessions with clean_start = false preserve state let session = SessionState::new("client1".to_string(), config.clone(), false); - // Add unacked messages for &id in &packet_ids { let packet = publish_packet(id, qos); session.store_unacked_publish(packet).await.unwrap(); @@ -145,8 +131,6 @@ mod clean_start_tests { let unacked_count = session.get_unacked_publishes().await.len(); prop_assert_eq!(unacked_count, packet_ids.len()); - // State persists across "reconnections" (new instance with same client_id) - // In practice, a session manager would handle this Ok(()) })?; } @@ -160,8 +144,7 @@ mod session_expiry_tests { proptest! { #[test] fn prop_session_expiry_interval_handling( - expiry in session_expiry(), - _elapsed_seconds in 0..7200u32 + expiry in session_expiry() ) { let rt = tokio::runtime::Runtime::new().unwrap(); rt.block_on(async { @@ -172,18 +155,8 @@ mod session_expiry_tests { let session = SessionState::new("client1".to_string(), config.clone(), false); - // For testing expiry, we'd need to manipulate last_activity time - // The is_expired() method checks time since last activity - // Session with expiry = 0 expires immediately on disconnect - if expiry == 0 { - // Would expire immediately after disconnect - prop_assert!(true); - } else if expiry == u32::MAX { - // Never expires - prop_assert!(!session.is_expired().await); - } else { - // Normal expiry interval + if expiry != 0 { prop_assert!(!session.is_expired().await); } Ok(()) @@ -199,7 +172,6 @@ mod session_expiry_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Add subscriptions for i in 0..sub_count { let topic_filter = format!("topic/{i}"); let sub = Subscription { @@ -209,14 +181,13 @@ mod session_expiry_tests { session.add_subscription(topic_filter, sub).await.unwrap(); } - // Queue messages for i in 0..msg_count { let msg = QueuedMessage { topic: format!("topic/{}", i % sub_count), - payload: vec![i as u8], + payload: vec![i.to_le_bytes()[0]], qos: QoS::AtLeastOnce, retain: false, - packet_id: Some((i as u16) + 1), + packet_id: Some(u16::try_from(i).unwrap() + 1), }; session.queue_message(msg).await.unwrap(); } @@ -246,7 +217,6 @@ mod subscription_management_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Add subscriptions for (topic, qos) in &topics { let topic_filter = topic.clone(); let sub = Subscription { @@ -259,20 +229,15 @@ mod subscription_management_tests { session.add_subscription(topic_filter, sub).await.unwrap(); } - // Check all subscriptions were added - // Note: MQTT replaces subscriptions to the same topic, so we need to count unique topics let all_subs = session.all_subscriptions().await; let unique_topics: std::collections::HashSet<_> = topics.iter().map(|(t, _)| t).collect(); prop_assert_eq!(all_subs.len(), unique_topics.len()); - // Check specific subscription retrieval - // Create a map of the final subscriptions (last one wins for duplicate topics) let mut expected_subs = std::collections::HashMap::new(); for (topic, qos) in &topics { expected_subs.insert(topic.clone(), *qos); } - // Verify each unique subscription for (topic, expected_qos) in expected_subs { let matching = session.matching_subscriptions(&topic).await; prop_assert!(!matching.is_empty(), "No subscription found for topic: {}", topic); @@ -292,7 +257,6 @@ mod subscription_management_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Add all subscriptions for topic in &topics { let topic_filter = topic.clone(); let sub = Subscription { @@ -304,7 +268,6 @@ mod subscription_management_tests { prop_assert_eq!(session.all_subscriptions().await.len(), topics.len()); - // Remove half of them let topics_vec: Vec<_> = topics.iter().cloned().collect(); let to_remove = topics_vec.len() / 2; for topic in &topics_vec[..to_remove] { @@ -314,7 +277,6 @@ mod subscription_management_tests { prop_assert_eq!(session.all_subscriptions().await.len(), topics.len() - to_remove); - // Verify correct ones remain for topic in &topics_vec[to_remove..] { let matching = session.matching_subscriptions(topic).await; prop_assert!(!matching.is_empty()); @@ -342,7 +304,6 @@ mod unacked_message_tests { let mut qos_packets = vec![]; let mut seen_ids = std::collections::HashSet::new(); - // Store unacked publishes, skipping duplicates (which would overwrite in real MQTT) for (id, qos) in packets { if qos != QoS::AtMostOnce && !seen_ids.contains(&id) { let packet = publish_packet(id, qos); @@ -352,11 +313,9 @@ mod unacked_message_tests { } } - // Verify all are tracked let unacked = session.get_unacked_publishes().await; prop_assert_eq!(unacked.len(), qos_packets.len()); - // Remove half and verify let half = qos_packets.len() / 2; for (id, _) in &qos_packets[..half] { let removed = session.remove_unacked_publish(*id).await; @@ -378,19 +337,16 @@ mod unacked_message_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Outbound QoS 2: we PUBLISH -> peer PUBRECs -> we PUBREL -> peer PUBCOMPs let packet = publish_packet(packet_id, QoS::ExactlyOnce); session.store_unacked_publish(packet).await.unwrap(); prop_assert!(!session.get_unacked_publishes().await.is_empty()); - // PUBREC received: the publish is retired and we owe a PUBREL session.complete_pubrec(packet_id).await; session.store_pubrel(packet_id).await; prop_assert!(session.get_unacked_publishes().await.is_empty()); prop_assert_eq!(session.get_unacked_pubrels().await.len(), 1); - // PUBCOMP received: flow complete session.complete_pubrel(packet_id).await; prop_assert!(session.get_unacked_pubrels().await.is_empty()); @@ -406,14 +362,11 @@ mod unacked_message_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Inbound QoS 2: peer PUBLISHes -> we PUBREC -> peer PUBRELs -> we PUBCOMP prop_assert!(session.mark_pubrec_pending(packet_id).await); prop_assert!(session.has_pubrec(packet_id).await); - // A redelivery of the same packet id is not a first receipt prop_assert!(!session.mark_pubrec_pending(packet_id).await); - // PUBREL received: flow complete, id released session.remove_pubrec(packet_id).await; prop_assert!(!session.has_pubrec(packet_id).await); @@ -429,8 +382,6 @@ mod unacked_message_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Inbound and outbound packet ids are independent namespaces that both - // allocate from 1, so the same id may be live in each direction at once. session.store_pubrel(packet_id).await; prop_assert!( @@ -450,7 +401,6 @@ mod unacked_message_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Store unacked pubrels, deduplicating let unique_ids: std::collections::HashSet<_> = packet_ids.iter().copied().collect(); for &id in &unique_ids { session.store_unacked_pubrel(id).await; @@ -459,7 +409,6 @@ mod unacked_message_tests { let pubrels = session.get_unacked_pubrels().await; prop_assert_eq!(pubrels.len(), unique_ids.len()); - // Remove half of unique IDs let unique_vec: Vec<_> = unique_ids.into_iter().collect(); let half = unique_vec.len() / 2; for &id in &unique_vec[..half] { @@ -495,33 +444,27 @@ mod message_queue_tests { let session = SessionState::new("client1".to_string(), config, true); - // Try to queue all messages for (i, size) in messages.iter().enumerate() { let msg = QueuedMessage { topic: format!("topic/{i}"), payload: vec![0u8; *size as usize], qos: QoS::AtLeastOnce, retain: false, - packet_id: Some((i as u16) + 1), + packet_id: Some(u16::try_from(i).unwrap() + 1), }; - // Always attempt to queue to test limits let _ = session.queue_message(msg).await; } - // Get the actual count from the session let actual_count = session.queued_message_count().await; - // Should not exceed max messages prop_assert!(actual_count <= 20, "Queued {} messages, exceeds limit of 20", actual_count); - // Verify we can dequeue some messages let to_dequeue = actual_count.min(5); let dequeued = session.dequeue_messages(to_dequeue).await; prop_assert!(dequeued.len() <= to_dequeue); - // Verify the count is updated after dequeuing let remaining = session.queued_message_count().await; prop_assert_eq!(remaining, actual_count - dequeued.len()); @@ -537,23 +480,19 @@ mod message_queue_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Queue messages for i in 0..count { let msg = QueuedMessage { topic: format!("topic/{i}"), - payload: vec![i as u8], + payload: vec![i.to_le_bytes()[0]], qos: QoS::AtLeastOnce, retain: false, - packet_id: Some((i as u16) + 1), + packet_id: Some(u16::try_from(i).unwrap() + 1), }; session.queue_message(msg).await.unwrap(); } prop_assert_eq!(session.queued_message_count().await, count); - // The SessionState queue_message method handles converting to ExpiringMessage - // We'll test that messages are queued correctly - // Message expiry is handled internally by the queue Ok(()) })?; @@ -575,32 +514,26 @@ mod flow_control_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Set receive maximum session.set_receive_maximum(receive_max).await; let mut in_flight: u16 = 0; - let mut _can_send_count = 0; - // Try to send messages for i in 0..message_count { if session.can_send_qos_message().await { - _can_send_count += 1; - let packet_id = (i as u16 % 65535) + 1; + let packet_id = u16::try_from(i).unwrap() + 1; if session.register_in_flight(packet_id).await.is_ok() { in_flight += 1; } } - // Randomly acknowledge some if i % 3 == 0 && i > 0 && in_flight > 0 { - let ack_id = ((i - 1) as u16 % 65535) + 1; + let ack_id = u16::try_from(i - 1).unwrap() + 1; if session.acknowledge_in_flight(ack_id).await.is_ok() { in_flight = in_flight.saturating_sub(1); } } } - // Should never exceed receive maximum prop_assert!(in_flight <= receive_max); Ok(()) @@ -616,13 +549,11 @@ mod flow_control_tests { rt.block_on(async { let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - // Set topic alias maximum session.set_topic_alias_maximum_out(max_alias).await; session.set_topic_alias_maximum_in(max_alias).await; let mut assigned_aliases = std::collections::HashMap::new(); - // Try to get aliases for topics for topic in &topics { if let Some(alias) = session.get_or_create_topic_alias(topic).await { prop_assert!(alias > 0 && alias <= max_alias); @@ -630,12 +561,10 @@ mod flow_control_tests { } } - // Should not exceed maximum prop_assert!(assigned_aliases.len() <= max_alias as usize); - // Test incoming alias registration for (i, topic) in topics.iter().take(max_alias as usize).enumerate() { - let alias = (i as u16) + 1; + let alias = u16::try_from(i).unwrap() + 1; session.register_incoming_topic_alias(alias, topic).await.unwrap(); let retrieved = session.get_topic_for_alias(alias).await; @@ -648,110 +577,6 @@ mod flow_control_tests { } } -#[cfg(test)] -mod retained_message_tests { - use super::*; - - proptest! { - #[test] - fn prop_retained_message_storage( - topics in prop::collection::vec("[a-zA-Z0-9/]{5,30}", 1..20) - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - - // Store retained messages - for (i, topic) in topics.iter().enumerate() { - let packet = PublishPacket { - dup: false, - qos: QoS::AtLeastOnce, - retain: true, - topic_name: topic.clone(), - packet_id: Some((i as u16) + 1), - properties: Properties::default(), - payload: vec![i as u8].into(), - protocol_version: 5, - stream_id: None, - }; - session.store_retained_message(&packet).await; - } - - // Retrieve by exact topic - for (i, topic) in topics.iter().enumerate() { - let retained = session.get_retained_messages(topic).await; - prop_assert_eq!(retained.len(), 1); - prop_assert_eq!(retained[0].payload[0], i as u8); - } - - // Clear retained message by sending empty payload - if let Some(topic) = topics.first() { - let clear_packet = PublishPacket { - dup: false, - qos: QoS::AtMostOnce, - retain: true, - topic_name: topic.clone(), - packet_id: None, - properties: Properties::default(), - payload: vec![].into(), - protocol_version: 5, - stream_id: None, - }; - session.store_retained_message(&clear_packet).await; - - let retained = session.get_retained_messages(topic).await; - prop_assert!(retained.is_empty()); - } - - Ok(()) - })?; - } - - #[test] - fn prop_retained_wildcard_matching( - base_topic in "[a-zA-Z0-9]{3,10}", - levels in 1..5usize - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let session = SessionState::new("client1".to_string(), SessionConfig::default(), true); - - // Create hierarchical topics - let mut topics = Vec::new(); - for i in 0..levels { - let topic = format!("{base_topic}/{i}/data"); - topics.push(topic.clone()); - - let packet = PublishPacket { - dup: false, - qos: QoS::AtLeastOnce, - retain: true, - topic_name: topic, - packet_id: Some((i as u16) + 1), - properties: Properties::default(), - payload: vec![i as u8].into(), - protocol_version: 5, - stream_id: None, - }; - session.store_retained_message(&packet).await; - } - - // Test wildcard matching - let filter = format!("{base_topic}/+/data"); - let matched = session.get_retained_messages(&filter).await; - prop_assert_eq!(matched.len(), levels); - - // Test multi-level wildcard - let filter = format!("{base_topic}/#"); - let matched = session.get_retained_messages(&filter).await; - prop_assert_eq!(matched.len(), levels); - - Ok(()) - })?; - } - } -} - #[cfg(test)] mod concurrent_session_tests { use super::*; @@ -773,7 +598,6 @@ mod concurrent_session_tests { let mut join_set = JoinSet::new(); - // Spawn concurrent tasks for thread_id in 0..thread_count { let session = Arc::clone(&session); let task = async move { @@ -789,10 +613,8 @@ mod concurrent_session_tests { join_set.spawn(task); } - // Wait for all tasks while join_set.join_next().await.is_some() {} - // Verify all subscriptions were added let all_subs = session.all_subscriptions().await; prop_assert_eq!(all_subs.len(), thread_count * ops_per_thread); @@ -820,7 +642,6 @@ mod concurrent_session_tests { let mut join_set = JoinSet::new(); - // Spawn concurrent tasks for thread_id in 0..thread_count { let session = Arc::clone(&session); let task = async move { @@ -828,10 +649,10 @@ mod concurrent_session_tests { for i in 0..msgs_per_thread { let msg = QueuedMessage { topic: format!("thread{thread_id}/msg{i}"), - payload: vec![thread_id as u8, i as u8], + payload: vec![u8::try_from(thread_id).unwrap(), u8::try_from(i).unwrap()], qos: QoS::AtLeastOnce, retain: false, - packet_id: Some(((thread_id * msgs_per_thread + i) as u16) + 1), + packet_id: Some(u16::try_from(thread_id * msgs_per_thread + i).unwrap() + 1), }; if session.queue_message(msg).await.is_ok() { queued += 1; @@ -842,13 +663,11 @@ mod concurrent_session_tests { join_set.spawn(task); } - // Collect results let mut total_queued = 0; while let Some(result) = join_set.join_next().await { total_queued += result.unwrap(); } - // Verify queued count let actual_count = session.queued_message_count().await; prop_assert_eq!(actual_count, total_queued); @@ -874,11 +693,9 @@ mod performance_property_tests { let start = Instant::now(); - // Perform mixed operations for i in 0..operation_count { match i % 5 { 0 => { - // Add subscription let topic_filter = format!("topic/{i}"); let sub = Subscription { topic_filter: topic_filter.clone(), @@ -889,24 +706,21 @@ mod performance_property_tests { 1 => { let msg = QueuedMessage { topic: format!("topic/{i}"), - payload: vec![i as u8], + payload: vec![i.to_le_bytes()[0]], qos: QoS::AtLeastOnce, retain: false, - packet_id: Some((i as u16) + 1), + packet_id: Some(u16::try_from(i).unwrap() + 1), }; let _ = session.queue_message(msg).await; } 2 => { - // Store unacked - let packet = publish_packet((i as u16) + 1, QoS::AtLeastOnce); + let packet = publish_packet(u16::try_from(i).unwrap() + 1, QoS::AtLeastOnce); let _ = session.store_unacked_publish(packet).await; } 3 => { - // Check stats let _ = session.stats().await; } 4 => { - // Touch session session.touch().await; } _ => unreachable!() @@ -914,9 +728,8 @@ mod performance_property_tests { } let elapsed = start.elapsed(); - let ops_per_second = operation_count as f64 / elapsed.as_secs_f64(); + let ops_per_second = f64::from(u32::try_from(operation_count).unwrap()) / elapsed.as_secs_f64(); - // Session operations should be reasonably fast prop_assert!( ops_per_second > 1_000.0, "Session operations too slow: {:.0} ops/sec", diff --git a/crates/mqtt5/tests/session_takeover_backlog.rs b/crates/mqtt5/tests/session_takeover_backlog.rs index dfcddf7c..2bd3cd32 100644 --- a/crates/mqtt5/tests/session_takeover_backlog.rs +++ b/crates/mqtt5/tests/session_takeover_backlog.rs @@ -41,6 +41,7 @@ fn qos1() -> SubscribeOptions { fn persistent(client_id: &str, clean_start: bool) -> ConnectOptions { ConnectOptions::new(client_id) .with_clean_start(clean_start) + .with_resume_existing_session(!clean_start) .with_session_expiry_interval(3600) .with_automatic_reconnect(false) } diff --git a/crates/mqtt5/tests/subscription_options_persistence.rs b/crates/mqtt5/tests/subscription_options_persistence.rs index c5f900ad..a58bdbde 100644 --- a/crates/mqtt5/tests/subscription_options_persistence.rs +++ b/crates/mqtt5/tests/subscription_options_persistence.rs @@ -1,5 +1,4 @@ #![cfg(feature = "broker")] -#![allow(clippy::large_futures)] mod common; use common::TestBroker; @@ -37,13 +36,16 @@ async fn test_no_local_persists_after_reconnect() { let received_clone = received.clone(); let client2 = MqttClient::new(client_id); - let session = client2 - .connect_with_options( + let session = Box::pin( + client2.connect_with_options( broker.address(), - ConnectOptions::new(client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); if !session.session_present { println!("Session not preserved, skipping test"); @@ -140,13 +142,16 @@ async fn test_retain_as_published_persists_after_reconnect() { let retain_clone = retain_flag_seen.clone(); let client2 = MqttClient::new(client_id); - let session = client2 - .connect_with_options( + let session = Box::pin( + client2.connect_with_options( broker.address(), - ConnectOptions::new(client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); if !session.session_present { println!("Session not preserved, skipping test"); @@ -227,13 +232,16 @@ async fn test_subscription_options_all_preserved() { let received_clone = received.clone(); let client2 = MqttClient::new(client_id); - let session = client2 - .connect_with_options( + let session = Box::pin( + client2.connect_with_options( broker.address(), - ConnectOptions::new(client_id).with_clean_start(false), - ) - .await - .unwrap(); + ConnectOptions::new(client_id) + .with_clean_start(false) + .with_resume_existing_session(true), + ), + ) + .await + .unwrap(); if !session.session_present { println!("Session not preserved, skipping test"); diff --git a/crates/mqttv5-cli/Cargo.toml b/crates/mqttv5-cli/Cargo.toml index b734c645..66b8161f 100644 --- a/crates/mqttv5-cli/Cargo.toml +++ b/crates/mqttv5-cli/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "mqttv5-cli" -version = "0.28.7" +version = "0.28.8" edition.workspace = true rust-version.workspace = true authors.workspace = true @@ -22,7 +22,7 @@ opentelemetry = ["mqtt5/opentelemetry"] codec = ["mqtt5/codec-all"] [dependencies] -mqtt5 = { path = "../mqtt5", version = "0.40" } +mqtt5 = { path = "../mqtt5", version = "0.41" } anyhow = "1.0.103" tracing = "0.1" serde = { version = "1.0", features = ["derive"] } diff --git a/crates/mqttv5-cli/src/commands/client_args.rs b/crates/mqttv5-cli/src/commands/client_args.rs new file mode 100644 index 00000000..01d1539c --- /dev/null +++ b/crates/mqttv5-cli/src/commands/client_args.rs @@ -0,0 +1,137 @@ +use clap::Args; +use mqtt5::{ProtocolVersion, QoS}; +use std::path::PathBuf; + +use super::parsers::{parse_duration_secs, parse_stream_strategy}; + +pub fn parse_qos(s: &str) -> Result { + match s { + "0" => Ok(QoS::AtMostOnce), + "1" => Ok(QoS::AtLeastOnce), + "2" => Ok(QoS::ExactlyOnce), + _ => Err(format!("QoS must be 0, 1, or 2, got: {s}")), + } +} + +pub fn parse_protocol_version(s: &str) -> Result { + match s { + "3.1.1" | "311" | "4" => Ok(ProtocolVersion::V311), + "5" | "5.0" => Ok(ProtocolVersion::V5), + _ => Err(format!("Invalid protocol version: {s}. Use '3.1.1' or '5'")), + } +} + +#[derive(Args)] +pub struct SessionArgs { + /// Don't clean start (resume existing session) + #[arg(long = "no-clean-start", env = "MQTT5_NO_CLEAN_START")] + pub no_clean_start: bool, + + /// Session expiry interval (e.g., 1h, 30m) (0 = expire on disconnect) + #[arg(long, value_parser = parse_duration_secs, env = "MQTT5_SESSION_EXPIRY")] + pub session_expiry: Option, + + /// Keep alive interval (e.g., 60s, 1m) (default: 60s) + #[arg(long, short = 'k', default_value = "60", value_parser = parse_duration_secs, env = "MQTT5_KEEP_ALIVE")] + pub keep_alive: u64, + + /// MQTT protocol version (3.1.1 or 5, default: 5) + #[arg(long, value_parser = parse_protocol_version, env = "MQTT5_PROTOCOL_VERSION")] + pub protocol_version: Option, +} + +#[derive(Args)] +pub struct WillArgs { + /// Will topic (last will and testament) + #[arg( + id = "will_topic", + long = "will-topic", + value_name = "WILL_TOPIC", + env = "MQTT5_WILL_TOPIC" + )] + pub topic: Option, + + /// Will message payload + #[arg( + id = "will_message", + long = "will-message", + value_name = "WILL_MESSAGE", + env = "MQTT5_WILL_MESSAGE" + )] + pub message: Option, + + /// Will `QoS` level (0, 1, or 2) + #[arg(id = "will_qos", long = "will-qos", value_name = "WILL_QOS", value_parser = parse_qos, env = "MQTT5_WILL_QOS")] + pub qos: Option, + + /// Will retain flag + #[arg(id = "will_retain", long = "will-retain", env = "MQTT5_WILL_RETAIN")] + pub retain: bool, +} + +#[derive(Args)] +pub struct TlsArgs { + /// TLS certificate file (PEM format) for secure connections + #[arg(long, env = "MQTT5_CERT")] + pub cert: Option, + + /// TLS private key file (PEM format) for secure connections + #[arg(long, env = "MQTT5_KEY")] + pub key: Option, + + /// TLS CA certificate file (PEM format) for server verification + #[arg(long, env = "MQTT5_CA_CERT")] + pub ca_cert: Option, + + /// Skip certificate verification for TLS/QUIC connections (insecure, for testing only) + #[arg(long, env = "MQTT5_INSECURE")] + pub insecure: bool, +} + +#[derive(Args)] +pub struct QuicArgs { + /// QUIC stream strategy (control-only, per-publish, per-topic, per-subscription) + #[arg(id = "quic_stream_strategy", long = "quic-stream-strategy", value_name = "QUIC_STREAM_STRATEGY", value_parser = parse_stream_strategy, env = "MQTT5_QUIC_STREAM_STRATEGY")] + pub stream_strategy: Option, + + /// Enable `MQoQ` flow headers for stream state tracking + #[arg( + id = "quic_flow_headers", + long = "quic-flow-headers", + env = "MQTT5_QUIC_FLOW_HEADERS" + )] + pub flow_headers: bool, + + /// Flow expiration interval (e.g., 5m, 1h) (default: 5m) + #[arg(id = "quic_flow_expire", long = "quic-flow-expire", value_name = "QUIC_FLOW_EXPIRE", default_value = "300", value_parser = parse_duration_secs, env = "MQTT5_QUIC_FLOW_EXPIRE")] + pub flow_expire: u64, + + /// Maximum concurrent QUIC streams + #[arg( + id = "quic_max_streams", + long = "quic-max-streams", + value_name = "QUIC_MAX_STREAMS", + env = "MQTT5_QUIC_MAX_STREAMS" + )] + pub max_streams: Option, + + /// Enable QUIC datagrams for unreliable transport + #[arg( + id = "quic_datagrams", + long = "quic-datagrams", + env = "MQTT5_QUIC_DATAGRAMS" + )] + pub datagrams: bool, + + /// QUIC connection timeout (e.g., 30s, 1m) (default: 30s) + #[arg(id = "quic_connect_timeout", long = "quic-connect-timeout", value_name = "QUIC_CONNECT_TIMEOUT", default_value = "30", value_parser = parse_duration_secs, env = "MQTT5_QUIC_CONNECT_TIMEOUT")] + pub connect_timeout: u64, + + /// Enable QUIC 0-RTT early data for faster reconnections + #[arg( + id = "quic_early_data", + long = "quic-early-data", + env = "MQTT5_QUIC_EARLY_DATA" + )] + pub early_data: bool, +} diff --git a/crates/mqttv5-cli/src/commands/mod.rs b/crates/mqttv5-cli/src/commands/mod.rs index e69e00df..d02b4120 100644 --- a/crates/mqttv5-cli/src/commands/mod.rs +++ b/crates/mqttv5-cli/src/commands/mod.rs @@ -1,6 +1,7 @@ pub mod acl_cmd; pub mod bench_cmd; pub mod broker_cmd; +pub mod client_args; pub mod parsers; pub mod passwd_cmd; pub mod pub_cmd; diff --git a/crates/mqttv5-cli/src/commands/pub_cmd.rs b/crates/mqttv5-cli/src/commands/pub_cmd.rs index 3d5e1387..a1a987ec 100644 --- a/crates/mqttv5-cli/src/commands/pub_cmd.rs +++ b/crates/mqttv5-cli/src/commands/pub_cmd.rs @@ -1,6 +1,3 @@ -#![allow(clippy::large_futures)] -#![allow(clippy::struct_excessive_bools)] - use anyhow::{Context, Result}; use clap::Args; use dialoguer::{Input, Select}; @@ -9,39 +6,24 @@ use mqtt5::time::Duration; #[cfg(feature = "codec")] use mqtt5::{CodecRegistry, DeflateCodec, GzipCodec}; use mqtt5::{ - ConnectOptions, ConnectionEvent, Message, MqttClient, ProtocolVersion, PublishOptions, QoS, - WillMessage, + ConnectOptions, ConnectionEvent, Message, MqttClient, PublishOptions, QoS, WillMessage, }; use std::io::{self, Read}; -use std::path::PathBuf; use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::Arc; use tokio::signal; use tokio::sync::Notify; use tracing::{debug, info, warn}; +use super::client_args::{parse_qos, QuicArgs, SessionArgs, TlsArgs, WillArgs}; use super::parsers::{ calculate_wait_until, duration_secs_to_u32, parse_duration_millis, parse_duration_secs, - parse_stream_strategy, }; #[derive(Args)] pub struct PubCommand { - /// MQTT topic to publish to - #[arg(long, short, env = "MQTT5_TOPIC")] - pub topic: Option, - - /// Message to publish - #[arg(long, short, env = "MQTT5_MESSAGE")] - pub message: Option, - - /// Read message from file - #[arg(long, short, env = "MQTT5_FILE")] - pub file: Option, - - /// Read message from stdin - #[arg(long, env = "MQTT5_STDIN")] - pub stdin: bool, + #[command(flatten)] + pub payload: PubPayloadArgs, /// Full broker URL for TLS/WebSocket/QUIC (e.g., , ) #[arg(long, short = 'U', conflicts_with_all = &["host", "port"], env = "MQTT5_URL")] @@ -55,45 +37,11 @@ pub struct PubCommand { #[arg(long, short, default_value = "1883", env = "MQTT5_PORT")] pub port: u16, - /// Quality of Service level (0, 1, or 2) - #[arg(long, short, value_parser = parse_qos, env = "MQTT5_QOS")] - pub qos: Option, - - /// Retain message - #[arg(long, short, env = "MQTT5_RETAIN")] - pub retain: bool, - - /// Message expiry interval in seconds (0 = no expiry) - #[arg(long, env = "MQTT5_MESSAGE_EXPIRY_INTERVAL")] - pub message_expiry_interval: Option, - - /// Topic alias (1-65535) for repeated publishing to same topic - #[arg(long, env = "MQTT5_TOPIC_ALIAS")] - pub topic_alias: Option, - - /// Response topic for request/response pattern (MQTT 5.0) - #[arg(long, env = "MQTT5_RESPONSE_TOPIC")] - pub response_topic: Option, - - /// Correlation data for request/response pattern (MQTT 5.0, hex-encoded) - #[arg(long, env = "MQTT5_CORRELATION_DATA")] - pub correlation_data: Option, - - /// Wait for response after publishing (requires --response-topic) - #[arg(long, env = "MQTT5_WAIT_RESPONSE")] - pub wait_response: bool, - - /// Timeout when waiting for response (e.g., 30s, 1m) (default: 30s) - #[arg(long, default_value = "30", value_parser = parse_duration_secs, env = "MQTT5_TIMEOUT")] - pub timeout: u64, + #[command(flatten)] + pub message_options: PubMessageOptionArgs, - /// Number of responses to wait for (default: 1, 0 = unlimited until timeout) - #[arg(long, default_value = "1", env = "MQTT5_RESPONSE_COUNT")] - pub response_count: u32, - - /// Output format for responses: raw, json, verbose - #[arg(long, default_value = "raw", value_parser = ["raw", "json", "verbose"], env = "MQTT5_OUTPUT_FORMAT")] - pub output_format: String, + #[command(flatten)] + pub response: PubResponseArgs, /// Username for authentication #[arg(long, short, env = "MQTT5_USERNAME")] @@ -119,53 +67,14 @@ pub struct PubCommand { #[arg(long, env = "MQTT5_NON_INTERACTIVE")] pub non_interactive: bool, - /// Don't clean start (resume existing session) - #[arg(long = "no-clean-start", env = "MQTT5_NO_CLEAN_START")] - pub no_clean_start: bool, - - /// Session expiry interval (e.g., 1h, 30m) (0 = expire on disconnect) - #[arg(long, value_parser = parse_duration_secs, env = "MQTT5_SESSION_EXPIRY")] - pub session_expiry: Option, - - /// Keep alive interval (e.g., 60s, 1m) (default: 60s) - #[arg(long, short = 'k', default_value = "60", value_parser = parse_duration_secs, env = "MQTT5_KEEP_ALIVE")] - pub keep_alive: u64, - - /// MQTT protocol version (3.1.1 or 5, default: 5) - #[arg(long, value_parser = parse_protocol_version, env = "MQTT5_PROTOCOL_VERSION")] - pub protocol_version: Option, - - /// Will topic (last will and testament) - #[arg(long, env = "MQTT5_WILL_TOPIC")] - pub will_topic: Option, - - /// Will message payload - #[arg(long, env = "MQTT5_WILL_MESSAGE")] - pub will_message: Option, - - /// Will `QoS` level (0, 1, or 2) - #[arg(long, value_parser = parse_qos, env = "MQTT5_WILL_QOS")] - pub will_qos: Option, - - /// Will retain flag - #[arg(long, env = "MQTT5_WILL_RETAIN")] - pub will_retain: bool, - - /// TLS certificate file (PEM format) for secure connections - #[arg(long, env = "MQTT5_CERT")] - pub cert: Option, + #[command(flatten)] + pub session: SessionArgs, - /// TLS private key file (PEM format) for secure connections - #[arg(long, env = "MQTT5_KEY")] - pub key: Option, + #[command(flatten)] + pub will: WillArgs, - /// TLS CA certificate file (PEM format) for server verification - #[arg(long, env = "MQTT5_CA_CERT")] - pub ca_cert: Option, - - /// Skip certificate verification for TLS/QUIC connections (insecure, for testing only) - #[arg(long, env = "MQTT5_INSECURE")] - pub insecure: bool, + #[command(flatten)] + pub tls: TlsArgs, /// Will delay interval (e.g., 5m, 1h) #[arg(long, value_parser = parse_duration_secs, env = "MQTT5_WILL_DELAY")] @@ -179,33 +88,8 @@ pub struct PubCommand { #[arg(long, env = "MQTT5_AUTO_RECONNECT")] pub auto_reconnect: bool, - /// QUIC stream strategy (control-only, per-publish, per-topic, per-subscription) - #[arg(long, value_parser = parse_stream_strategy, env = "MQTT5_QUIC_STREAM_STRATEGY")] - pub quic_stream_strategy: Option, - - /// Enable `MQoQ` flow headers for stream state tracking - #[arg(long, env = "MQTT5_QUIC_FLOW_HEADERS")] - pub quic_flow_headers: bool, - - /// Flow expiration interval (e.g., 5m, 1h) (default: 5m) - #[arg(long, default_value = "300", value_parser = parse_duration_secs, env = "MQTT5_QUIC_FLOW_EXPIRE")] - pub quic_flow_expire: u64, - - /// Maximum concurrent QUIC streams - #[arg(long, env = "MQTT5_QUIC_MAX_STREAMS")] - pub quic_max_streams: Option, - - /// Enable QUIC datagrams for unreliable transport - #[arg(long, env = "MQTT5_QUIC_DATAGRAMS")] - pub quic_datagrams: bool, - - /// QUIC connection timeout (e.g., 30s, 1m) (default: 30s) - #[arg(long, default_value = "30", value_parser = parse_duration_secs, env = "MQTT5_QUIC_CONNECT_TIMEOUT")] - pub quic_connect_timeout: u64, - - /// Enable QUIC 0-RTT early data for faster reconnections - #[arg(long, env = "MQTT5_QUIC_EARLY_DATA")] - pub quic_early_data: bool, + #[command(flatten)] + pub quic: QuicArgs, /// Delay before publishing (e.g., 5s, 1m30s) #[arg(long, value_parser = parse_duration_secs, env = "MQTT5_DELAY")] @@ -264,33 +148,81 @@ pub struct PubCommand { pub codec_min_size: usize, } -fn parse_qos(s: &str) -> Result { - match s { - "0" => Ok(QoS::AtMostOnce), - "1" => Ok(QoS::AtLeastOnce), - "2" => Ok(QoS::ExactlyOnce), - _ => Err(format!("QoS must be 0, 1, or 2, got: {s}")), - } +#[derive(Args)] +pub struct PubPayloadArgs { + /// MQTT topic to publish to + #[arg(long, short, env = "MQTT5_TOPIC")] + pub topic: Option, + + /// Message to publish + #[arg(long, short, env = "MQTT5_MESSAGE")] + pub message: Option, + + /// Read message from file + #[arg(long, short, env = "MQTT5_FILE")] + pub file: Option, + + /// Read message from stdin + #[arg(long, env = "MQTT5_STDIN")] + pub stdin: bool, } -fn parse_protocol_version(s: &str) -> Result { - match s { - "3.1.1" | "311" | "4" => Ok(ProtocolVersion::V311), - "5" | "5.0" => Ok(ProtocolVersion::V5), - _ => Err(format!("Invalid protocol version: {s}. Use '3.1.1' or '5'")), - } +#[derive(Args)] +pub struct PubMessageOptionArgs { + /// Quality of Service level (0, 1, or 2) + #[arg(long, short, value_parser = parse_qos, env = "MQTT5_QOS")] + pub qos: Option, + + /// Retain message + #[arg(long, short, env = "MQTT5_RETAIN")] + pub retain: bool, + + /// Message expiry interval in seconds (0 = no expiry) + #[arg(long, env = "MQTT5_MESSAGE_EXPIRY_INTERVAL")] + pub message_expiry_interval: Option, + + /// Topic alias (1-65535) for repeated publishing to same topic + #[arg(long, env = "MQTT5_TOPIC_ALIAS")] + pub topic_alias: Option, +} + +#[derive(Args)] +pub struct PubResponseArgs { + /// Response topic for request/response pattern (MQTT 5.0) + #[arg(long, env = "MQTT5_RESPONSE_TOPIC")] + pub response_topic: Option, + + /// Correlation data for request/response pattern (MQTT 5.0, hex-encoded) + #[arg(long, env = "MQTT5_CORRELATION_DATA")] + pub correlation_data: Option, + + /// Wait for response after publishing (requires --response-topic) + #[arg(long, env = "MQTT5_WAIT_RESPONSE")] + pub wait_response: bool, + + /// Timeout when waiting for response (e.g., 30s, 1m) (default: 30s) + #[arg(long, default_value = "30", value_parser = parse_duration_secs, env = "MQTT5_TIMEOUT")] + pub timeout: u64, + + /// Number of responses to wait for (default: 1, 0 = unlimited until timeout) + #[arg(long, default_value = "1", env = "MQTT5_RESPONSE_COUNT")] + pub response_count: u32, + + /// Output format for responses: raw, json, verbose + #[arg(long, default_value = "raw", value_parser = ["raw", "json", "verbose"], env = "MQTT5_OUTPUT_FORMAT")] + pub output_format: String, } fn prompt_topic_and_qos(cmd: &mut PubCommand) -> Result<(String, QoS)> { - if cmd.topic.is_none() && !cmd.non_interactive { + if cmd.payload.topic.is_none() && !cmd.non_interactive { let topic = Input::::new() .with_prompt("MQTT topic (e.g., sensors/temperature, home/status)") .interact() .context("Failed to get topic input")?; - cmd.topic = Some(topic); + cmd.payload.topic = Some(topic); } - let topic = cmd.topic.take().ok_or_else(|| { + let topic = cmd.payload.topic.take().ok_or_else(|| { anyhow::anyhow!("Topic is required. Use --topic or run without --non-interactive") })?; @@ -312,11 +244,11 @@ fn prompt_topic_and_qos(cmd: &mut PubCommand) -> Result<(String, QoS)> { ); } - if cmd.wait_response && cmd.response_topic.is_none() { + if cmd.response.wait_response && cmd.response.response_topic.is_none() { anyhow::bail!("--response-topic is required when using --wait-response"); } - let qos = if cmd.qos.is_none() && !cmd.non_interactive { + let qos = if cmd.message_options.qos.is_none() && !cmd.non_interactive { let qos_options = vec![ "0 (At most once - fire and forget)", "1 (At least once - acknowledged)", @@ -335,10 +267,10 @@ fn prompt_topic_and_qos(cmd: &mut PubCommand) -> Result<(String, QoS)> { _ => QoS::AtMostOnce, } } else { - cmd.qos.unwrap_or(QoS::AtMostOnce) + cmd.message_options.qos.unwrap_or(QoS::AtMostOnce) }; - if cmd.wait_response && qos == QoS::AtMostOnce { + if cmd.response.wait_response && qos == QoS::AtMostOnce { warn!("Using --wait-response with QoS 0 may be unreliable; consider using -q 1 or -q 2"); } @@ -347,18 +279,19 @@ fn prompt_topic_and_qos(cmd: &mut PubCommand) -> Result<(String, QoS)> { fn build_connect_options(cmd: &PubCommand, client_id: &str) -> ConnectOptions { let mut options = ConnectOptions::new(client_id.to_owned()) - .with_clean_start(!cmd.no_clean_start) - .with_keep_alive(Duration::from_secs(cmd.keep_alive)); + .with_clean_start(!cmd.session.no_clean_start) + .with_resume_existing_session(cmd.session.no_clean_start) + .with_keep_alive(Duration::from_secs(cmd.session.keep_alive)); if cmd.auto_reconnect { options = options.with_automatic_reconnect(true); } - if let Some(version) = cmd.protocol_version { + if let Some(version) = cmd.session.protocol_version { options = options.with_protocol_version(version); } - if let Some(expiry) = cmd.session_expiry { + if let Some(expiry) = cmd.session.session_expiry { options = options.with_session_expiry_interval(duration_secs_to_u32(expiry)); } @@ -411,11 +344,11 @@ async fn configure_auth( } fn configure_will(options: &mut ConnectOptions, cmd: &PubCommand) { - if let Some(topic) = cmd.will_topic.clone() { - let payload = cmd.will_message.clone().unwrap_or_default(); - let mut will = WillMessage::new(topic, payload.into_bytes()).with_retain(cmd.will_retain); + if let Some(topic) = cmd.will.topic.clone() { + let payload = cmd.will.message.clone().unwrap_or_default(); + let mut will = WillMessage::new(topic, payload.into_bytes()).with_retain(cmd.will.retain); - if let Some(qos) = cmd.will_qos { + if let Some(qos) = cmd.will.qos { will = will.with_qos(qos); } @@ -428,29 +361,29 @@ fn configure_will(options: &mut ConnectOptions, cmd: &PubCommand) { } async fn configure_quic_transport(client: &MqttClient, cmd: &PubCommand) { - if let Some(strategy) = cmd.quic_stream_strategy { + if let Some(strategy) = cmd.quic.stream_strategy { client.set_quic_stream_strategy(strategy).await; debug!("QUIC stream strategy: {:?}", strategy); } - if cmd.quic_flow_headers { + if cmd.quic.flow_headers { client.set_quic_flow_headers(true).await; debug!("QUIC flow headers enabled"); } client - .set_quic_flow_expire(std::time::Duration::from_secs(cmd.quic_flow_expire)) + .set_quic_flow_expire(std::time::Duration::from_secs(cmd.quic.flow_expire)) .await; - if let Some(max) = cmd.quic_max_streams { + if let Some(max) = cmd.quic.max_streams { client.set_quic_max_streams(Some(max)).await; debug!("QUIC max streams: {max}"); } - if cmd.quic_datagrams { + if cmd.quic.datagrams { client.set_quic_datagrams(true).await; debug!("QUIC datagrams enabled"); } client - .set_quic_connect_timeout(Duration::from_secs(cmd.quic_connect_timeout)) + .set_quic_connect_timeout(Duration::from_secs(cmd.quic.connect_timeout)) .await; - if cmd.quic_early_data { + if cmd.quic.early_data { client.set_quic_early_data(true).await; debug!("QUIC 0-RTT early data enabled"); } @@ -464,17 +397,17 @@ async fn configure_tls_certs( let is_secure = broker_url.starts_with("ssl://") || broker_url.starts_with("mqtts://") || broker_url.starts_with("quics://"); - let has_certs = cmd.cert.is_some() || cmd.key.is_some() || cmd.ca_cert.is_some(); + let has_certs = cmd.tls.cert.is_some() || cmd.tls.key.is_some() || cmd.tls.ca_cert.is_some(); if is_secure && has_certs { - let cert_pem = if let Some(cert_path) = &cmd.cert { + let cert_pem = if let Some(cert_path) = &cmd.tls.cert { Some(std::fs::read(cert_path).with_context(|| { format!("Failed to read certificate file: {}", cert_path.display()) })?) } else { None }; - let key_pem = if let Some(key_path) = &cmd.key { + let key_pem = if let Some(key_path) = &cmd.tls.key { Some( std::fs::read(key_path) .with_context(|| format!("Failed to read key file: {}", key_path.display()))?, @@ -482,7 +415,7 @@ async fn configure_tls_certs( } else { None }; - let ca_pem = if let Some(ca_path) = &cmd.ca_cert { + let ca_pem = if let Some(ca_path) = &cmd.tls.ca_cert { Some(std::fs::read(ca_path).with_context(|| { format!("Failed to read CA certificate file: {}", ca_path.display()) })?) @@ -495,6 +428,13 @@ async fn configure_tls_certs( Ok(()) } +fn awaited_response_topic(cmd: &PubCommand) -> Option<&str> { + cmd.response + .response_topic + .as_deref() + .filter(|_| cmd.response.wait_response) +} + async fn setup_response_subscription( client: &MqttClient, cmd: &PubCommand, @@ -503,16 +443,15 @@ async fn setup_response_subscription( let received_count = Arc::new(AtomicU32::new(0)); let done_notify = Arc::new(Notify::new()); - if cmd.wait_response { - let response_topic = cmd.response_topic.as_ref().unwrap().clone(); + if let Some(response_topic) = awaited_response_topic(cmd) { let expected_correlation = correlation_data; - let target_count = cmd.response_count; - let output_format = cmd.output_format.clone(); + let target_count = cmd.response.response_count; + let output_format = cmd.response.output_format.clone(); let received_clone = received_count.clone(); let done_clone = done_notify.clone(); client - .subscribe(&response_topic, move |msg: Message| { + .subscribe(response_topic, move |msg: Message| { if let Some(ref expected) = expected_correlation { match &msg.properties.correlation_data { Some(received) if received == expected => {} @@ -529,10 +468,7 @@ async fn setup_response_subscription( }) .await?; - debug!( - "SUBACK received for '{}', subscription ready before publish", - response_topic - ); + debug!("SUBACK received for '{response_topic}', subscription ready before publish"); } Ok((received_count, done_notify)) @@ -572,10 +508,10 @@ async fn publish_loop( qos: QoS, correlation_data: Option<&Vec>, ) -> Result<()> { - let has_properties = cmd.retain - || cmd.message_expiry_interval.is_some() - || cmd.topic_alias.is_some() - || cmd.response_topic.is_some() + let has_properties = cmd.message_options.retain + || cmd.message_options.message_expiry_interval.is_some() + || cmd.message_options.topic_alias.is_some() + || cmd.response.response_topic.is_some() || correlation_data.is_some(); let repeat_count = cmd.repeat.unwrap_or(1); @@ -621,15 +557,15 @@ async fn publish_with_properties( ) -> Result<()> { let mut options = PublishOptions { qos, - retain: cmd.retain, + retain: cmd.message_options.retain, ..Default::default() }; - options.properties.message_expiry_interval = cmd.message_expiry_interval; - options.properties.topic_alias = cmd.topic_alias; + options.properties.message_expiry_interval = cmd.message_options.message_expiry_interval; + options.properties.topic_alias = cmd.message_options.topic_alias; options .properties .response_topic - .clone_from(&cmd.response_topic); + .clone_from(&cmd.response.response_topic); options.properties.correlation_data = correlation_data.cloned(); client .publish_with_options(topic, message.as_bytes(), options) @@ -661,7 +597,7 @@ fn print_publish_result(cmd: &PubCommand, iteration: u64, topic: &str, qos: QoS) } else { println!("✓ Published message to '{}' (QoS {})", topic, qos as u8); } - if cmd.retain { + if cmd.message_options.retain { println!(" Message retained on broker"); } } @@ -672,17 +608,14 @@ async fn wait_for_response( received_count: Arc, client: &MqttClient, ) -> Result<()> { - if !cmd.wait_response { + let Some(response_topic) = awaited_response_topic(cmd) else { return Ok(()); - } + }; - let timeout_secs = cmd.timeout; - let target_count = cmd.response_count; + let timeout_secs = cmd.response.timeout; + let target_count = cmd.response.response_count; - println!( - "Waiting for response on '{}'...", - cmd.response_topic.as_ref().unwrap() - ); + println!("Waiting for response on '{response_topic}'..."); tokio::select! { () = done_notify.notified() => { @@ -799,7 +732,7 @@ pub async fn execute(mut cmd: PubCommand, verbose: bool, debug: bool) -> Result< #[cfg(feature = "codec")] configure_codec(&mut options, &cmd)?; - if cmd.insecure { + if cmd.tls.insecure { client.set_insecure_tls(true).await; info!("Insecure TLS mode enabled (certificate verification disabled)"); } @@ -819,17 +752,18 @@ pub async fn execute(mut cmd: PubCommand, verbose: bool, debug: bool) -> Result< info!("Publishing to topic '{}'...", topic); - if cmd.topic_alias == Some(0) { + if cmd.message_options.topic_alias == Some(0) { anyhow::bail!("Topic alias must be between 1 and 65535, got: 0"); } - let correlation_data: Option> = if cmd.wait_response && cmd.correlation_data.is_none() { - Some(format!("rr-{}", rand::rng().random::()).into_bytes()) - } else if let Some(ref hex_data) = cmd.correlation_data { - Some(hex::decode(hex_data).context("Invalid hex in --correlation-data")?) - } else { - None - }; + let correlation_data: Option> = + if cmd.response.wait_response && cmd.response.correlation_data.is_none() { + Some(format!("rr-{}", rand::rng().random::()).into_bytes()) + } else if let Some(ref hex_data) = cmd.response.correlation_data { + Some(hex::decode(hex_data).context("Invalid hex in --correlation-data")?) + } else { + None + }; let (received_count, done_notify) = setup_response_subscription(&client, &cmd, correlation_data.clone()).await?; @@ -915,7 +849,7 @@ fn configure_codec(options: &mut ConnectOptions, cmd: &PubCommand) -> Result<()> } async fn get_message_content(cmd: &mut PubCommand) -> Result { - if cmd.stdin { + if cmd.payload.stdin { debug!("Reading message from stdin"); let mut buffer = String::new(); io::stdin() @@ -924,7 +858,7 @@ async fn get_message_content(cmd: &mut PubCommand) -> Result { return Ok(buffer.trim().to_string()); } - if let Some(file_path) = &cmd.file { + if let Some(file_path) = &cmd.payload.file { debug!("Reading message from file: {}", file_path); let content = tokio::fs::read_to_string(file_path) .await @@ -932,7 +866,7 @@ async fn get_message_content(cmd: &mut PubCommand) -> Result { return Ok(content.trim().to_string()); } - if let Some(message) = &cmd.message { + if let Some(message) = &cmd.payload.message { return Ok(message.clone()); } diff --git a/crates/mqttv5-cli/src/commands/sub_cmd.rs b/crates/mqttv5-cli/src/commands/sub_cmd.rs index b3b94e62..956c37ee 100644 --- a/crates/mqttv5-cli/src/commands/sub_cmd.rs +++ b/crates/mqttv5-cli/src/commands/sub_cmd.rs @@ -1,6 +1,3 @@ -#![allow(clippy::large_futures)] -#![allow(clippy::struct_excessive_bools)] - use anyhow::{Context, Result}; use clap::Args; use dialoguer::{Input, Select}; @@ -8,15 +5,15 @@ use mqtt5::client::auth_handlers::{JwtAuthHandler, ScramSha256AuthHandler}; use mqtt5::time::Duration; #[cfg(feature = "codec")] use mqtt5::{CodecRegistry, DeflateCodec, GzipCodec}; -use mqtt5::{ConnectOptions, MqttClient, ProtocolVersion, QoS, WillMessage}; -use std::path::PathBuf; +use mqtt5::{ConnectOptions, MqttClient, QoS, WillMessage}; use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::Arc; use tokio::signal; use tokio::sync::Notify; use tracing::{debug, info}; -use super::parsers::{duration_secs_to_u32, parse_duration_secs, parse_stream_strategy}; +use super::client_args::{parse_qos, QuicArgs, SessionArgs, TlsArgs, WillArgs}; +use super::parsers::{duration_secs_to_u32, parse_duration_secs}; #[derive(Args)] pub struct SubCommand { @@ -60,121 +57,31 @@ pub struct SubCommand { #[arg(long, short, env = "MQTT5_CLIENT_ID")] pub client_id: Option, - /// Print verbose output (include topic names) - #[arg(long, short)] - pub verbose: bool, - - /// Print message properties (`QoS`, retain, expiry, content type, response topic, user properties, subscription IDs) - #[arg(long = "show-properties", short = 's', env = "MQTT5_SHOW_PROPERTIES")] - pub show_properties: bool, - - /// Skip prompts and use defaults/fail if required args missing - #[arg(long, env = "MQTT5_NON_INTERACTIVE")] - pub non_interactive: bool, - - /// Number of messages to receive before exiting (0 = infinite) - #[arg(long, short = 'n', default_value = "0", env = "MQTT5_COUNT")] - pub count: u32, - - /// Don't clean start (resume existing session) - #[arg(long = "no-clean-start", env = "MQTT5_NO_CLEAN_START")] - pub no_clean_start: bool, - - /// Session expiry interval (e.g., 1h, 30m) (0 = expire on disconnect) - #[arg(long, value_parser = parse_duration_secs, env = "MQTT5_SESSION_EXPIRY")] - pub session_expiry: Option, - - /// Keep alive interval (e.g., 60s, 1m) (default: 60s) - #[arg(long, short = 'k', default_value = "60", value_parser = parse_duration_secs, env = "MQTT5_KEEP_ALIVE")] - pub keep_alive: u64, - - /// MQTT protocol version (3.1.1 or 5, default: 5) - #[arg(long, value_parser = parse_protocol_version, env = "MQTT5_PROTOCOL_VERSION")] - pub protocol_version: Option, - - /// Will topic (last will and testament) - #[arg(long, env = "MQTT5_WILL_TOPIC")] - pub will_topic: Option, + #[command(flatten)] + pub output: SubOutputArgs, - /// Will message payload - #[arg(long, env = "MQTT5_WILL_MESSAGE")] - pub will_message: Option, + #[command(flatten)] + pub session: SessionArgs, - /// Will `QoS` level (0, 1, or 2) - #[arg(long, value_parser = parse_qos, env = "MQTT5_WILL_QOS")] - pub will_qos: Option, - - /// Will retain flag - #[arg(long, env = "MQTT5_WILL_RETAIN")] - pub will_retain: bool, + #[command(flatten)] + pub will: WillArgs, /// Will delay interval (e.g., 5m, 1h) #[arg(long, value_parser = parse_duration_secs, env = "MQTT5_WILL_DELAY")] pub will_delay: Option, - /// TLS certificate file (PEM format) for secure connections - #[arg(long, env = "MQTT5_CERT")] - pub cert: Option, - - /// TLS private key file (PEM format) for secure connections - #[arg(long, env = "MQTT5_KEY")] - pub key: Option, - - /// TLS CA certificate file (PEM format) for server verification - #[arg(long, env = "MQTT5_CA_CERT")] - pub ca_cert: Option, - - /// Skip certificate verification for TLS/QUIC connections (insecure, for testing only) - #[arg(long, env = "MQTT5_INSECURE")] - pub insecure: bool, + #[command(flatten)] + pub tls: TlsArgs, /// Enable automatic reconnection when broker disconnects #[arg(long, env = "MQTT5_AUTO_RECONNECT")] pub auto_reconnect: bool, - /// No Local - if true, Application Messages published by this client will not be received back - #[arg(long, env = "MQTT5_NO_LOCAL")] - pub no_local: bool, - - /// Subscription identifier (1-268435455) to identify which subscription matched a message - #[arg(long, env = "MQTT5_SUBSCRIPTION_IDENTIFIER")] - pub subscription_identifier: Option, - - /// Retain handling: 0=SendAtSubscribe, 1=SendAtSubscribeIfNew, 2=DoNotSend - #[arg(long, value_parser = parse_retain_handling, env = "MQTT5_RETAIN_HANDLING")] - pub retain_handling: Option, - - /// Retain As Published - keep original retain flag when delivering messages - #[arg(long, env = "MQTT5_RETAIN_AS_PUBLISHED")] - pub retain_as_published: bool, - - /// QUIC stream strategy (control-only, per-publish, per-topic, per-subscription) - #[arg(long, value_parser = parse_stream_strategy, env = "MQTT5_QUIC_STREAM_STRATEGY")] - pub quic_stream_strategy: Option, + #[command(flatten)] + pub subscription_options: SubscriptionOptionArgs, - /// Enable `MQoQ` flow headers for stream state tracking - #[arg(long, env = "MQTT5_QUIC_FLOW_HEADERS")] - pub quic_flow_headers: bool, - - /// Flow expiration interval (e.g., 5m, 1h) (default: 5m) - #[arg(long, default_value = "300", value_parser = parse_duration_secs, env = "MQTT5_QUIC_FLOW_EXPIRE")] - pub quic_flow_expire: u64, - - /// Maximum concurrent QUIC streams - #[arg(long, env = "MQTT5_QUIC_MAX_STREAMS")] - pub quic_max_streams: Option, - - /// Enable QUIC datagrams for unreliable transport - #[arg(long, env = "MQTT5_QUIC_DATAGRAMS")] - pub quic_datagrams: bool, - - /// QUIC connection timeout (e.g., 30s, 1m) (default: 30s) - #[arg(long, default_value = "30", value_parser = parse_duration_secs, env = "MQTT5_QUIC_CONNECT_TIMEOUT")] - pub quic_connect_timeout: u64, - - /// Enable QUIC 0-RTT early data for faster reconnections - #[arg(long, env = "MQTT5_QUIC_EARLY_DATA")] - pub quic_early_data: bool, + #[command(flatten)] + pub quic: QuicArgs, /// OpenTelemetry OTLP endpoint (e.g., `http://localhost:4317`) #[cfg(feature = "opentelemetry")] @@ -197,21 +104,42 @@ pub struct SubCommand { pub codec: Option, } -fn parse_qos(s: &str) -> Result { - match s { - "0" => Ok(QoS::AtMostOnce), - "1" => Ok(QoS::AtLeastOnce), - "2" => Ok(QoS::ExactlyOnce), - _ => Err(format!("QoS must be 0, 1, or 2, got: {s}")), - } +#[derive(Args)] +pub struct SubOutputArgs { + /// Print verbose output (include topic names) + #[arg(long, short)] + pub verbose: bool, + + /// Print message properties (`QoS`, retain, expiry, content type, response topic, user properties, subscription IDs) + #[arg(long = "show-properties", short = 's', env = "MQTT5_SHOW_PROPERTIES")] + pub show_properties: bool, + + /// Skip prompts and use defaults/fail if required args missing + #[arg(long, env = "MQTT5_NON_INTERACTIVE")] + pub non_interactive: bool, + + /// Number of messages to receive before exiting (0 = infinite) + #[arg(long, short = 'n', default_value = "0", env = "MQTT5_COUNT")] + pub count: u32, } -fn parse_protocol_version(s: &str) -> Result { - match s { - "3.1.1" | "311" | "4" => Ok(ProtocolVersion::V311), - "5" | "5.0" => Ok(ProtocolVersion::V5), - _ => Err(format!("Invalid protocol version: {s}. Use '3.1.1' or '5'")), - } +#[derive(Args)] +pub struct SubscriptionOptionArgs { + /// No Local - if true, Application Messages published by this client will not be received back + #[arg(long, env = "MQTT5_NO_LOCAL")] + pub no_local: bool, + + /// Subscription identifier (1-268435455) to identify which subscription matched a message + #[arg(long, env = "MQTT5_SUBSCRIPTION_IDENTIFIER")] + pub subscription_identifier: Option, + + /// Retain handling: 0=SendAtSubscribe, 1=SendAtSubscribeIfNew, 2=DoNotSend + #[arg(long, value_parser = parse_retain_handling, env = "MQTT5_RETAIN_HANDLING")] + pub retain_handling: Option, + + /// Retain As Published - keep original retain flag when delivering messages + #[arg(long, env = "MQTT5_RETAIN_AS_PUBLISHED")] + pub retain_as_published: bool, } fn parse_retain_handling(s: &str) -> Result { @@ -224,7 +152,7 @@ fn parse_retain_handling(s: &str) -> Result { } fn prompt_topic_and_qos(cmd: &mut SubCommand) -> Result<(String, QoS)> { - if cmd.topic.is_none() && !cmd.non_interactive { + if cmd.topic.is_none() && !cmd.output.non_interactive { let topic = Input::::new() .with_prompt("MQTT topic to subscribe to (e.g., sensors/+, home/#)") .interact() @@ -238,7 +166,7 @@ fn prompt_topic_and_qos(cmd: &mut SubCommand) -> Result<(String, QoS)> { validate_topic_filter(&topic)?; - let qos = if cmd.qos.is_none() && !cmd.non_interactive { + let qos = if cmd.qos.is_none() && !cmd.output.non_interactive { let qos_options = vec!["0 (At most once)", "1 (At least once)", "2 (Exactly once)"]; let selection = Select::new() .with_prompt("Quality of Service level") @@ -261,18 +189,19 @@ fn prompt_topic_and_qos(cmd: &mut SubCommand) -> Result<(String, QoS)> { fn build_connect_options(cmd: &SubCommand, client_id: &str) -> ConnectOptions { let mut options = ConnectOptions::new(client_id.to_owned()) - .with_clean_start(!cmd.no_clean_start) - .with_keep_alive(Duration::from_secs(cmd.keep_alive)); + .with_clean_start(!cmd.session.no_clean_start) + .with_resume_existing_session(cmd.session.no_clean_start) + .with_keep_alive(Duration::from_secs(cmd.session.keep_alive)); if cmd.auto_reconnect { options = options.with_automatic_reconnect(true); } - if let Some(version) = cmd.protocol_version { + if let Some(version) = cmd.session.protocol_version { options = options.with_protocol_version(version); } - if let Some(expiry) = cmd.session_expiry { + if let Some(expiry) = cmd.session.session_expiry { options = options.with_session_expiry_interval(duration_secs_to_u32(expiry)); } @@ -325,11 +254,11 @@ async fn configure_auth( } fn configure_will(options: &mut ConnectOptions, cmd: &SubCommand) { - if let Some(topic) = cmd.will_topic.clone() { - let payload = cmd.will_message.clone().unwrap_or_default(); - let mut will = WillMessage::new(topic, payload.into_bytes()).with_retain(cmd.will_retain); + if let Some(topic) = cmd.will.topic.clone() { + let payload = cmd.will.message.clone().unwrap_or_default(); + let mut will = WillMessage::new(topic, payload.into_bytes()).with_retain(cmd.will.retain); - if let Some(qos) = cmd.will_qos { + if let Some(qos) = cmd.will.qos { will = will.with_qos(qos); } @@ -342,29 +271,29 @@ fn configure_will(options: &mut ConnectOptions, cmd: &SubCommand) { } async fn configure_quic_transport(client: &MqttClient, cmd: &SubCommand) { - if let Some(strategy) = cmd.quic_stream_strategy { + if let Some(strategy) = cmd.quic.stream_strategy { client.set_quic_stream_strategy(strategy).await; debug!("QUIC stream strategy: {:?}", strategy); } - if cmd.quic_flow_headers { + if cmd.quic.flow_headers { client.set_quic_flow_headers(true).await; debug!("QUIC flow headers enabled"); } client - .set_quic_flow_expire(std::time::Duration::from_secs(cmd.quic_flow_expire)) + .set_quic_flow_expire(std::time::Duration::from_secs(cmd.quic.flow_expire)) .await; - if let Some(max) = cmd.quic_max_streams { + if let Some(max) = cmd.quic.max_streams { client.set_quic_max_streams(Some(max)).await; debug!("QUIC max streams: {max}"); } - if cmd.quic_datagrams { + if cmd.quic.datagrams { client.set_quic_datagrams(true).await; debug!("QUIC datagrams enabled"); } client - .set_quic_connect_timeout(Duration::from_secs(cmd.quic_connect_timeout)) + .set_quic_connect_timeout(Duration::from_secs(cmd.quic.connect_timeout)) .await; - if cmd.quic_early_data { + if cmd.quic.early_data { client.set_quic_early_data(true).await; debug!("QUIC 0-RTT early data enabled"); } @@ -379,19 +308,20 @@ async fn configure_tls_certs( || broker_url.starts_with("mqtts://") || broker_url.starts_with("quics://"); - if !is_secure || (cmd.cert.is_none() && cmd.key.is_none() && cmd.ca_cert.is_none()) { + if !is_secure || (cmd.tls.cert.is_none() && cmd.tls.key.is_none() && cmd.tls.ca_cert.is_none()) + { return Ok(()); } let cert_pem = - if let Some(cert_path) = &cmd.cert { + if let Some(cert_path) = &cmd.tls.cert { Some(std::fs::read(cert_path).with_context(|| { format!("Failed to read certificate file: {}", cert_path.display()) })?) } else { None }; - let key_pem = if let Some(key_path) = &cmd.key { + let key_pem = if let Some(key_path) = &cmd.tls.key { Some( std::fs::read(key_path) .with_context(|| format!("Failed to read key file: {}", key_path.display()))?, @@ -399,7 +329,7 @@ async fn configure_tls_certs( } else { None }; - let ca_pem = if let Some(ca_path) = &cmd.ca_cert { + let ca_pem = if let Some(ca_path) = &cmd.tls.ca_cert { Some(std::fs::read(ca_path).with_context(|| { format!("Failed to read CA certificate file: {}", ca_path.display()) })?) @@ -485,9 +415,9 @@ async fn subscribe_and_print( cmd: &SubCommand, qos: QoS, ) -> Result> { - let target_count = cmd.count; - let verbose = cmd.verbose; - let show_properties = cmd.show_properties; + let target_count = cmd.output.count; + let verbose = cmd.output.verbose; + let show_properties = cmd.output.show_properties; info!("Subscribing to '{}' (QoS {})...", topic, qos as u8); @@ -495,7 +425,7 @@ async fn subscribe_and_print( let done_notify = Arc::new(Notify::new()); let done_notify_clone = done_notify.clone(); - if let Some(sub_id) = cmd.subscription_identifier { + if let Some(sub_id) = cmd.subscription_options.subscription_identifier { if sub_id == 0 || sub_id > 268_435_455 { anyhow::bail!("Subscription identifier must be between 1 and 268435455, got: {sub_id}"); } @@ -503,12 +433,13 @@ async fn subscribe_and_print( let subscribe_options = mqtt5::SubscribeOptions { qos, - no_local: cmd.no_local, - retain_as_published: cmd.retain_as_published, + no_local: cmd.subscription_options.no_local, + retain_as_published: cmd.subscription_options.retain_as_published, retain_handling: cmd + .subscription_options .retain_handling .unwrap_or(mqtt5::RetainHandling::SendAtSubscribe), - subscription_identifier: cmd.subscription_identifier, + subscription_identifier: cmd.subscription_options.subscription_identifier, }; let (packet_id, granted_qos) = client @@ -615,7 +546,7 @@ pub async fn execute(mut cmd: SubCommand, verbose: bool, debug: bool) -> Result< #[cfg(feature = "codec")] configure_codec(&mut options, &cmd)?; - if cmd.insecure { + if cmd.tls.insecure { client.set_insecure_tls(true).await; info!("Insecure TLS mode enabled (certificate verification disabled)"); } From 0ead001a9a87d62786fb0e0cd189c59f87c40d77 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Fabr=C3=ADcio=20Bracht?= Date: Wed, 23 Sep 2026 22:29:44 -0300 Subject: [PATCH 2/2] fix quorum review findings and report publish outcomes --- .github/workflows/rust.yml | 22 + CHANGELOG.md | 61 +- Makefile.toml | 30 +- WASM_USAGE.md | 40 +- crates/mqtt5-conformance/CONFORMANCE_DIARY.md | 12 +- .../src/test_client/inprocess.rs | 17 +- crates/mqtt5-wasm/Cargo.toml | 2 +- crates/mqtt5-wasm/README.md | 2 +- crates/mqtt5-wasm/src/client/callbacks.rs | 39 +- crates/mqtt5-wasm/src/client/connection.rs | 19 +- crates/mqtt5-wasm/src/client/handlers.rs | 12 +- crates/mqtt5-wasm/src/client/mod.rs | 88 +- crates/mqtt5-wasm/src/client/outbound.rs | 95 +- crates/mqtt5-wasm/src/client/qos.rs | 120 +- crates/mqtt5-wasm/src/client/reconnect.rs | 12 +- crates/mqtt5-wasm/src/client/state.rs | 47 +- .../mqtt5-wasm/src/client_handler/connect.rs | 21 +- .../src/client_handler/subscribe.rs | 20 +- crates/mqtt5-wasm/tests/client_broker.rs | 101 ++ crates/mqtt5-wasm/tests/conformance_client.rs | 678 +++++++- crates/mqtt5/README.md | 32 +- crates/mqtt5/src/broker/bridge/connection.rs | 9 +- .../src/broker/client_handler/connect.rs | 24 +- .../src/broker/client_handler/subscribe.rs | 26 +- crates/mqtt5/src/broker/router.rs | 602 +++---- crates/mqtt5/src/client/direct/ack.rs | 124 +- crates/mqtt5/src/client/direct/handlers.rs | 11 +- crates/mqtt5/src/client/direct/keepalive.rs | 35 +- crates/mqtt5/src/client/direct/mod.rs | 1494 +++++++++++++++-- crates/mqtt5/src/client/direct/outbound.rs | 56 +- crates/mqtt5/src/client/direct/reader.rs | 277 ++- crates/mqtt5/src/client/direct/replay.rs | 349 +++- crates/mqtt5/src/client/direct/tracking.rs | 316 ++++ crates/mqtt5/src/client/direct/unified.rs | 62 +- crates/mqtt5/src/client/mock.rs | 39 +- crates/mqtt5/src/client/mod.rs | 73 +- crates/mqtt5/src/client/publish_outcome.rs | 215 +++ crates/mqtt5/src/client/trait.rs | 5 +- crates/mqtt5/src/lib.rs | 9 +- crates/mqtt5/src/session/flow_control.rs | 41 +- crates/mqtt5/src/session/state.rs | 31 + crates/mqtt5/src/types.rs | 15 +- .../mqtt5/tests/broker_bridge_integration.rs | 18 +- crates/mqtt5/tests/change_only_delivery.rs | 114 +- crates/mqtt5/tests/client_publish.rs | 34 +- crates/mqtt5/tests/conf_client_a.rs | 2 +- crates/mqtt5/tests/conf_client_b.rs | 79 +- crates/mqtt5/tests/conf_client_c.rs | 4 + crates/mqtt5/tests/conf_client_d.rs | 4 +- .../mqtt5/tests/conf_client_offline_queue.rs | 973 +++++++++++ crates/mqtt5/tests/conf_client_quic.rs | 360 ++++ .../mqtt5/tests/integration_complete_flow.rs | 24 +- crates/mqtt5/tests/integration_no_local.rs | 130 +- .../tests/integration_retain_as_published.rs | 30 +- crates/mqtt5/tests/message_queuing.rs | 66 +- crates/mqtt5/tests/mock_client.rs | 40 +- .../mqtt5/tests/outbound_receive_maximum.rs | 4 +- crates/mqtt5/tests/qos_flow.rs | 86 +- crates/mqtt5/tests/server_max_packet_size.rs | 8 +- .../mqtt5/tests/shared_subscription_basic.rs | 14 +- crates/mqtt5/tests/turmoil_multi_client.rs | 123 +- crates/mqtt5/tests/turmoil_pubsub.rs | 104 +- .../tests/turmoil_shared_subscriptions.rs | 16 +- crates/mqttv5-cli/src/commands/bench_cmd.rs | 4 +- crates/mqttv5-cli/src/commands/pub_cmd.rs | 37 +- specs/tla/offline-queue/OfflineQueue.cfg | 37 + specs/tla/offline-queue/OfflineQueue.tla | 523 ++++++ .../OfflineQueue_D0_ExactlyOnce.cfg | 23 + .../OfflineQueue_D0_MaximumPacketSize.cfg | 23 + .../OfflineQueue_D0_MaximumQoS.cfg | 23 + .../OfflineQueue_D0_NoSilentLoss.cfg | 23 + .../OfflineQueue_D0_NoStaleServerPid.cfg | 23 + .../offline-queue/OfflineQueue_D0_Order.cfg | 23 + .../OfflineQueue_D0_PidUnique.cfg | 23 + .../OfflineQueue_D0_QoSFidelity.cfg | 23 + .../OfflineQueue_D0_RetainAvailable.cfg | 23 + .../offline-queue/OfflineQueue_D0_live.cfg | 25 + .../OfflineQueue_NEG_alldelivered.cfg | 25 + .../OfflineQueue_NEG_impossible.cfg | 25 + .../OfflineQueue_NEG_noconnectfair.cfg | 25 + .../OfflineQueue_NEG_quarantine.cfg | 25 + .../OfflineQueue_V_noquarantine.cfg | 23 + .../offline-queue/OfflineQueue_V_pidevent.cfg | 23 + .../OfflineQueue_V_replaydowngrade.cfg | 37 + .../offline-queue/OfflineQueue_alldims.cfg | 37 + specs/tla/offline-queue/OfflineQueue_live.cfg | 26 + specs/tla/offline-queue/OfflineQueue_qos.cfg | 37 + specs/tla/offline-queue/OfflineQueue_rm2.cfg | 37 + specs/tla/offline-queue/OfflineQueue_size.cfg | 37 + specs/tla/offline-queue/README.md | 294 ++++ 90 files changed, 7368 insertions(+), 1632 deletions(-) create mode 100644 crates/mqtt5-wasm/tests/client_broker.rs create mode 100644 crates/mqtt5/src/client/direct/tracking.rs create mode 100644 crates/mqtt5/src/client/publish_outcome.rs create mode 100644 crates/mqtt5/tests/conf_client_offline_queue.rs create mode 100644 crates/mqtt5/tests/conf_client_quic.rs create mode 100644 specs/tla/offline-queue/OfflineQueue.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue.tla create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_ExactlyOnce.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_MaximumPacketSize.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_MaximumQoS.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_NoSilentLoss.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_NoStaleServerPid.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_Order.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_PidUnique.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_QoSFidelity.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_RetainAvailable.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_D0_live.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_NEG_alldelivered.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_NEG_impossible.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_NEG_noconnectfair.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_NEG_quarantine.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_V_noquarantine.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_V_pidevent.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_V_replaydowngrade.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_alldims.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_live.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_qos.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_rm2.cfg create mode 100644 specs/tla/offline-queue/OfflineQueue_size.cfg create mode 100644 specs/tla/offline-queue/README.md diff --git a/.github/workflows/rust.yml b/.github/workflows/rust.yml index 860b5c65..ddbb2624 100644 --- a/.github/workflows/rust.yml +++ b/.github/workflows/rust.yml @@ -167,6 +167,28 @@ jobs: - name: Check WASM compilation run: cargo check --target wasm32-unknown-unknown --all-features -p mqtt5-wasm + - name: Install Node.js + uses: actions/setup-node@v4 + with: + node-version: 22 + + - name: Install wasm-bindgen-test-runner + run: | + version=$(awk '/^name = "wasm-bindgen"$/ { getline; gsub(/"/, "", $3); print $3; exit }' Cargo.lock) + test -n "$version" + installed=$(wasm-bindgen-test-runner --version 2>/dev/null || true) + if [ "$installed" != "wasm-bindgen-test-runner $version" ]; then + cargo install wasm-bindgen-cli --version "$version" --locked + fi + + - name: Run WASM tests + env: + CARGO_TARGET_WASM32_UNKNOWN_UNKNOWN_RUNNER: wasm-bindgen-test-runner + WASM_BINDGEN_TEST_TIMEOUT: "120" + run: | + cargo test -p mqtt5-wasm --target wasm32-unknown-unknown + cargo test -p mqtt5-wasm --target wasm32-unknown-unknown --features broker + - name: Build WASM package run: | cd crates/mqtt5-wasm diff --git a/CHANGELOG.md b/CHANGELOG.md index 24e75251..f6d721af 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -21,6 +21,19 @@ A client-side conformance audit drove the real `MqttClient` against a raw-byte f - `test_utils::test_retained_message` - `test_utils::TestMessageBuilder::build_retained_batch` - **Removed the deprecated `WebSocketConfig::with_tls_verification` and `WebSocketConfig::verify_tls`.** Nothing read `verify_tls`. Use `tls_config`. +- **`PublishResult` now reports the real outcome of a publish.** + - It is `Sent(Delivery)` for an acknowledged publish, or `Queued(PublishHandle)` for a QoS 1/2 publish whose outcome is still open. That covers one queued while offline, and a live publish whose connection ended, whose client disconnected, or whose 10 s acknowledgement wait elapsed before the ack arrived. Such a publish used to return `Err` even though the message stayed in the session and was resent. `publish()` now returns an error only when the message definitely was not delivered. `QoS0` and `QoS1Or2 { packet_id }` are removed. + - `Delivery` says which QoS was actually used: `Unconfirmed`, `AtLeastOnce { packet_id }` or `ExactlyOnce { packet_id }`. + - A `PublishHandle` can be awaited. It settles exactly once to a `PublishOutcome`: + - `Delivered(Delivery)` + - `Rejected(PublishRejection)`: definitely not delivered + - `Indeterminate(IndeterminateReason)`: may have been delivered + - `PublishResult::outcome()` gives the outcome on either path. + - `mqtt5` no longer re-exports `mqtt5_protocol::PublishResult`. + - The design is verified in TLA+ in `specs/tla/offline-queue/`. +- **`MessageRouter::subscribe` and `subscribe_as` now take a `SubscriptionRequest`** instead of 10 or 11 separate arguments. Build one with `SubscriptionRequest::new(client_id, topic_filter, qos)` and the `with_*` setters. The defaults match the values callers passed before. +- **The broker rejects packets with non-zero reserved fixed-header flags on every packet type** (`[MQTT-2.1.3-1]`). This covers CONNECT, PINGREQ, DISCONNECT and AUTH, which were previously accepted; SUBSCRIBE, UNSUBSCRIBE, PUBREL and the acks were already checked. The check comes from `mqtt5-protocol` 0.15.2, so an `mqtt5` 0.40 broker picks it up through `cargo update`. +- **With deferred ack enabled, an unresolved `AckToken` holds back every later PUBACK/PUBREC on that connection**, including automatic acks for plain subscriptions, because acks must go out in arrival order (`[MQTT-4.6.0-2]`, `[MQTT-4.6.0-3]`). An application that holds a token until some later message arrives can stall itself once the broker's in-flight window fills. ### Fixed @@ -37,6 +50,30 @@ A client-side conformance audit drove the real `MqttClient` against a raw-byte f - Request Problem Information=0: a Reason String or User Property on a packet other than PUBLISH, CONNACK or DISCONNECT is a protocol error - **Acknowledgements go out in PUBLISH arrival order when deferred ack is enabled** (`[MQTT-4.6.0-2]`, `[MQTT-4.6.0-3]`). Automatic acks and `AckToken` acks now share one ordered release, so a later ack waits for earlier pending ones. `AckToken::reject` maps reason codes that are invalid for PUBACK/PUBREC to 0x80. Acks still queued when the connection drops are discarded if the new connection reports Session Present=0. - **WebSocket reads reassemble MQTT packets from the byte stream** (`[MQTT-6.0.0-2]`). Several packets in one frame, or one packet split across frames, used to corrupt payloads or drop the session. A text frame closes the connection (`[MQTT-6.0.0-1]`), and a WebSocket Ping no longer ends the session. +- **QUIC connections close on protocol errors too.** After sending DISCONNECT the client closes the QUIC connection with an application close code and stops its stream readers. Before, it kept accepting server streams and acknowledging messages after its own DISCONNECT. The client's Maximum Packet Size is enforced on QUIC control and data streams (0x95). Topic Aliases are resolved on unidirectional streams. A malformed or oversized packet on a data stream now fails the connection instead of silently dropping that stream. The same goes for a QoS 1/2 PUBLISH on a unidirectional stream, which cannot carry its acknowledgement (DISCONNECT 0x82). Before, it was dropped silently. +- **A publish that waits for send quota across a reconnect is checked again against the new connection** before it is sent: Maximum QoS, Retain Available, Maximum Packet Size and Topic Alias range. Its quota claim is bound to the connection it was taken on, so it can no longer go out uncounted and exceed Receive Maximum (`[MQTT-3.3.4-7]`). +- **Acks released before a Session Present=0 reconnect are dropped** instead of being applied to the new session. Before, a later QoS 2 message reusing that packet identifier could be suppressed as a duplicate. +- **Offline-queued messages are no longer lost silently.** + - An offline RETAIN publish is rejected immediately if the last CONNACK reported Retain Available=0. + - At flush, a queued message that no longer fits the new connection (RETAIN not available, larger than Maximum Packet Size) is reported `Rejected` and the flush continues. It used to be dropped with only a log line after `publish()` had returned success. + - A message above the new Maximum QoS is downgraded, and its outcome reports the QoS actually used. +- **Session resume no longer resends messages the new connection does not allow.** Before resending, each stored PUBLISH is checked against the new Retain Available, Maximum Packet Size and Maximum QoS. One that fails is not sent and is reported `Indeterminate`. A QoS 2 packet identifier abandoned this way is kept out of reuse until the session is lost, so the broker cannot treat a new message as a duplicate. PUBRELs are always resent. +- **Unacknowledged messages are handled explicitly when the session is lost or discarded.** When a Clean Start=0 reconnect gets Session Present=0, unacknowledged QoS 1 messages are sent again first, in order. A QoS 2 message still waiting for PUBREC is reported `Indeterminate`. One that already received PUBREC Success is reported `Delivered`, because the receiver owns it from that point. A Clean Start=1 connect discards unacknowledged session state and reports it (`[MQTT-3.1.2-4]`). Queued messages that were never sent are kept. +- **Races in the offline flush are closed:** + - A flush or replay task from a replaced connection can no longer write to that connection or store after reconnect. + - The packet identifiers of queued, in-flush and staged messages cannot be reallocated. + - Send quota is released when storing a flushed message fails. +- **Connection loss closes the send quota.** Publishes waiting for quota fail immediately with `NotConnected`, instead of waiting out the 30 s backpressure timeout. An interrupted offline flush no longer keeps handles pending after the client is dropped; they resolve `Indeterminate(Abandoned)`. +- **Acknowledgements settle the publish outcome before session state is released**, so a disconnect while an ack is being processed cannot leave a delivered publish unsettled or misreported. +- **A QoS 2 packet identifier abandoned during replay is quarantined before it leaves the session store**, and a queued message downgraded to QoS 0 stays queued if the connection ends before it is written. +- **A replay stuck on a dead, replaced connection can no longer block the next connection's replay and flush or starve its send quota.** A replay task stops writing as soon as its connection ends. +- **Connection loss detected by keepalive fully ends the connection.** It closes the send quota, releases in-flight publish waiters, and stops the packet reader. +- **A PUBACK, PUBREC or PUBCOMP that does not match the QoS stage of its packet identifier is a protocol error** (DISCONNECT 0x82). Before, it released the outbound state. New: `session::state::OutboundStage` and `SessionState::outbound_stage`. +- **PUBACK, PUBREC and PUBCOMP received on a QUIC server-opened data stream settle the publish outcome and are stage-checked** exactly as on the control stream. +- **A PUBREC with an error reason code reports the publish as rejected only after its stored state has been removed.** If processing is interrupted before that, the publish stays pending and is resent on session resume. An error-code PUBREC for a packet identifier already at the PUBREL stage is a protocol error (DISCONNECT 0x82). +- **The broker bridge counts a publish as sent only once it is acknowledged.** +- **Packet identifier allocation no longer scans the offline queue.** With tens of thousands of queued messages, each `publish` used to spend over a second of CPU without yielding. +- **On MQTT 3.1.1 connections, v5-only publish properties are ignored instead of rejected.** They were never encoded anyway. This covers Topic Alias, Response Topic and Subscription Identifier. Topic names are still validated. - **CONNECT carries `request_problem_information`, `request_response_information` and user properties**, which were silently dropped. The client never sends AUTH when CONNECT had no Authentication Method (`[MQTT-4.12.0-7]`). The Assigned Client Identifier is adopted for later reconnects (`[MQTT-3.1.3-2]`). ## [mqttv5-cli 0.28.8] - 2026-09-23 @@ -44,16 +81,36 @@ A client-side conformance audit drove the real `MqttClient` against a raw-byte f ### Changed - Depends on `mqtt5` 0.41. `pub` and `sub` with `--no-clean-start` set `resume_existing_session`, so they resume the broker-held session as before. +- `pub` and `bench` report an error when a QoS 1/2 publish is not acknowledged, instead of printing success. -## [mqtt5-wasm 1.5.0] - 2026-09-23 +## [mqtt5-wasm 2.0.0] - 2026-09-23 + +### Breaking + +- **The Rust option methods are now snake_case.** On `WasmConnectOptions`, `WasmReconnectOptions`, `WasmPublishOptions`, `WasmSubscribeOptions`, `WasmWillMessage` and `MessageProperties`, `set_keepAlive` is now `set_keep_alive`, `cleanStart` is now `clean_start`, and so on. Rust code that calls these methods must be updated. **The JavaScript/TypeScript API is unchanged**: every property and method keeps its camelCase JS name. Only the raw wasm export symbols in the generated `InitOutput` interface are renamed (for example `connectoptions_cleanStart` is now `connectoptions_clean_start`). That affects only code that calls the raw exports directly. +- **A freshly created client now rejects Session Present=1** with DISCONNECT 0x82, and `connect` fails, as `[MQTT-3.2.2-4]` requires. Set `resumeExistingSession` to resume a broker-held session from a new client instance on purpose. A client instance that has connected before, including one that auto-reconnects, resumes normally. +- **Invalid requests are rejected before anything is sent**, and QoS 1/2 publishes wait while the server's Receive Maximum is exhausted. +- **A publish above the server's Maximum QoS is downgraded and the QoS used is reported.** `publishWithOptions` now resolves with the QoS used (`Promise`, previously `Promise`), and `publishQos1`/`publishQos2` callbacks receive it as a second argument. At Maximum QoS 0 they resolve with packet id 0 and call the callback with `(0, 0)`. +- **Pending publishes settle with an `indeterminate: ...` error, not reason code 128, when their outcome is unknown.** This happens when the session is discarded by Clean Start, lost (Session Present=0), ends with the connection, or holds a message that no longer fits the new connection's limits. The message says the publish may have been delivered. ### Added - **`ConnectOptions.resumeExistingSession`.** It lets a freshly created client accept Session Present=1 from a broker-held session. The `session-recovery` and `qos2-recovery` examples use it. +### Fixed + +- **The browser client now follows the same client-side conformance rules as the native client.** This release fixes the missing PUBACK for inbound QoS 1 messages, byte loss when several packets arrived in one frame, packet identifier reuse, and the missing resend on session resume. It also adds topic and filter validation, and enforcement of the server's Receive Maximum, Topic Alias Maximum, Maximum QoS, Retain Available and Maximum Packet Size. Protocol errors now send DISCONNECT with a reason code and close the transport. With `keepAlive=0`, the client no longer sends PINGREQs. Tests: `crates/mqtt5-wasm/tests/conformance_client.rs`, now run in CI under Node. +- **Session lifetime follows the protocol.** With MQTT 3.1.1 and `cleanStart=false`, the session now survives connection loss. Before, it was discarded, and the reconnect was then rejected forever. With MQTT v5, the Session Expiry Interval from the server's CONNACK takes precedence over the requested one. Automatic reconnects send Clean Start=0 only while the client still holds session state (or `resumeExistingSession` is set), and Clean Start=1 otherwise. The native client instead always reconnects with the configured Clean Start. +- **`disconnect()` settles pending publish promises and QoS callbacks** with a "message remains in session" error when the session outlives the connection. Before, they could hang forever. A later `connect` on the same instance resumes the session and resends. +- **Resent PUBRELs are counted against a lowered Receive Maximum** after a resume (`[MQTT-3.3.4-7]`). +- **Session resume re-checks each unacknowledged PUBLISH against the new CONNACK** (Retain Available, Maximum Packet Size, Maximum QoS). A message that no longer fits is not resent, and its promise or callback settles as indeterminate. A QoS 2 packet identifier abandoned this way is not reused until the session is discarded, so the broker cannot mistake a new message for a duplicate. PUBRELs are still resent. +- **When a Clean Start=0 reconnect gets Session Present=0**, unacknowledged QoS 1 publishes are sent again as new messages (DUP=0, original order) and resolve on acknowledgement. Unacknowledged QoS 2 publishes settle as indeterminate. A connect with Clean Start=1 discards the client's session state before CONNECT (`[MQTT-3.1.2-4]`). +- **A publish waiting for send quota fails instead of going out on a different connection** if the connection changed while it waited. +- **Publishes on MQTT 3.1.1 connections are encoded as 3.1.1.** They used to be encoded as v5, which corrupted the payload. v5-only properties are ignored on 3.1.1. +- **The in-browser broker no longer panics on the first routed PUBLISH.** Before, `tokio::time::Instant` was called on wasm32. + ### Changed -- **The browser client now follows the same client-side conformance rules as the native client.** This release fixes the missing PUBACK for inbound QoS 1 messages, byte loss when several packets arrived in one frame, packet identifier reuse, and the missing resend on session resume. It also adds topic and filter validation, and enforcement of the server's Receive Maximum, Topic Alias Maximum, Maximum QoS, Retain Available and Maximum Packet Size. Protocol errors now send DISCONNECT with a reason code and close the transport. With `keepAlive=0`, the client no longer sends PINGREQs. A fresh client now rejects Session Present=1 unless `resumeExistingSession` is set. Invalid requests are rejected before sending, and QoS 1/2 publishes wait while the server's Receive Maximum is exhausted. Tests: `crates/mqtt5-wasm/tests/conformance_client.rs`. - Depends on `mqtt5` 0.41 and `mqtt5-protocol` 0.15.2. ## [mqtt5-protocol 0.15.2] - 2026-09-23 diff --git a/Makefile.toml b/Makefile.toml index fbcd5741..d4972ace 100644 --- a/Makefile.toml +++ b/Makefile.toml @@ -43,7 +43,7 @@ echo "" echo "🌐 WASM BUILD (mqtt5-wasm crate)" echo " cargo make wasm-check Check WASM compilation" echo " cargo make wasm-clippy Run clippy for WASM target" -echo " cargo make wasm-test Run WASM crate tests (native)" +echo " cargo make wasm-test Run WASM crate tests under Node" echo " cargo make wasm-build Build WASM package with wasm-pack" echo " cargo make wasm-verify Run all WASM checks" echo " cargo make wasm-examples Build WASM and copy to examples" @@ -252,10 +252,32 @@ command = "cargo" args = ["clippy", "--target", "wasm32-unknown-unknown", "-p", "mqtt5-wasm", "--features", "broker", "--", "-D", "warnings", "-W", "clippy::pedantic"] description = "Run clippy for WASM target (strict, pedantic)" +[tasks.wasm-test-runner] +script = ''' +#!/usr/bin/env bash +set -euo pipefail +version=$(awk '/^name = "wasm-bindgen"$/ { getline; gsub(/"/, "", $3); print $3; exit }' Cargo.lock) +if [ -z "$version" ]; then + echo "wasm-bindgen not found in Cargo.lock" >&2 + exit 1 +fi +installed=$(wasm-bindgen-test-runner --version 2>/dev/null || true) +if [ "$installed" != "wasm-bindgen-test-runner $version" ]; then + cargo install wasm-bindgen-cli --version "$version" --locked +fi +''' +description = "Install the wasm-bindgen-test-runner matching the Cargo.lock wasm-bindgen version" + [tasks.wasm-test] -command = "cargo" -args = ["test", "-p", "mqtt5-wasm"] -description = "Run WASM crate tests (native)" +dependencies = ["wasm-test-runner"] +env = { CARGO_TARGET_WASM32_UNKNOWN_UNKNOWN_RUNNER = "wasm-bindgen-test-runner", WASM_BINDGEN_TEST_TIMEOUT = "120" } +script = ''' +#!/usr/bin/env bash +set -euo pipefail +cargo test -p mqtt5-wasm --target wasm32-unknown-unknown +cargo test -p mqtt5-wasm --target wasm32-unknown-unknown --features broker +''' +description = "Run WASM crate tests under Node with wasm-bindgen-test-runner" [tasks.wasm-verify] dependencies = ["wasm-check", "wasm-clippy", "wasm-test"] diff --git a/WASM_USAGE.md b/WASM_USAGE.md index e7304a3e..0270817a 100644 --- a/WASM_USAGE.md +++ b/WASM_USAGE.md @@ -262,15 +262,23 @@ client.destroy(); ```javascript await client.publish(topic, payloadBytes); -await client.publishWithOptions(topic, payloadBytes, publishOptions); +const qosUsed = await client.publishWithOptions(topic, payloadBytes, publishOptions); +// resolves with the QoS the message was sent at (0, 1 or 2) const packetId = await client.publishQos1(topic, payloadBytes, callback); -// callback(reasonCode) called when PUBACK received +// callback(reasonCode, qosUsed) called when PUBACK received const packetId = await client.publishQos2(topic, payloadBytes, callback); -// callback(reasonCode) called when PUBCOMP received +// callback(reasonCode, qosUsed) called when PUBCOMP received ``` +A publish never goes out above the server's Maximum QoS (from CONNACK). A higher requested QoS is downgraded to the server maximum and the QoS actually used is reported: `publishWithOptions` resolves with it, and the `publishQos1`/`publishQos2` callback receives it as the second argument. When the server maximum is 0, the message is sent at QoS 0, `publishQos1`/`publishQos2` resolve with packet identifier `0` (none) and the callback is called immediately with `(0, 0)`. Retain Available and Maximum Packet Size are not adjusted: a publish that violates them is rejected before anything is sent. + +`disconnect()` ends this instance's use of the connection and never leaves an acknowledgement pending: + +- If the session ends with the connection (see [Session Lifetime](#session-lifetime)), the session state is discarded: pending `publishWithOptions` promises reject with `Publish not acknowledged: indeterminate: session ended with the connection; may have been delivered` and `publishQos1`/`publishQos2` callbacks receive that `indeterminate: ...` string. +- If the session outlives the connection, the session state is kept. Pending `publishWithOptions` promises reject with `Publish not acknowledged: disconnected; message remains in session and will be resent on resume by this client instance`, and `publishQos1`/`publishQos2` callbacks receive that `disconnected; ...` string. A later `connect*()` with `cleanStart = false` on the same `MqttClient` instance resumes the session (when the server reports Session Present = 1) and resends those messages; their acknowledgements are then handled silently. + #### Subscribing ```javascript @@ -572,6 +580,8 @@ const opts = new ConnectOptions(); opts.keepAlive = 60; // default: 60 seconds opts.cleanStart = true; // default: true +opts.resumeExistingSession = false; // default: false; with cleanStart = false, accept + // Session Present = 1 without local session state opts.username = 'alice'; // default: null opts.set_password(encoder.encode('pw')); // accepts Uint8Array opts.protocolVersion = 5; // 4 (v3.1.1) or 5 (v5.0), default: 5 @@ -605,6 +615,30 @@ opts.clearCodecRegistry(); await client.connectWithOptions('ws://broker:8000/mqtt', opts); ``` +#### Session Lifetime + +The client keeps its session state (unacknowledged QoS 1/2 publishes and QoS 2 releases) for as long as the session outlives the network connection: + +- MQTT v5.0: while the Session Expiry Interval is greater than 0. The value the server returns in CONNACK replaces the requested `sessionExpiryInterval`. +- MQTT v3.1.1 (`protocolVersion = 4`): while `cleanStart = false` (CleanSession = 0). With `cleanStart = true` the session ends with the connection. + +When the session ends with the connection, pending publishes are settled on connection loss or `disconnect()` with `indeterminate: session ended with the connection; may have been delivered`. A `connect*()` with `cleanStart = true` discards any session state the instance still holds before sending CONNECT (MQTT-3.1.2-4): nothing is resent, pending publishes settle with `indeterminate: session discarded by clean start; may have been delivered`, and quarantined packet identifiers are released. When it outlives the connection, a connection loss leaves them pending and they are resent when the session resumes (automatic reconnect, or a later `connect*()` with `cleanStart = false` on the same `MqttClient` instance). Automatic reconnects send `cleanStart = false` only while the client holds session state (or `resumeExistingSession` is set), so a v3.1.1 `cleanStart = true` client never turns into a persistent session. + +On the next CONNACK after a `cleanStart = false` CONNECT on the same `MqttClient` instance, the unacknowledged messages are handled according to Session Present and the limits of the new connection (Retain Available, Maximum Packet Size, Maximum QoS): + +| CONNACK | Unacknowledged message | Outcome | +|---|---|---| +| Session Present = 1 | QoS 1/2 PUBLISH within the new limits | Resent with the same packet identifier, in the original order (DUP = 1 when it was already sent on this session) | +| Session Present = 1 | QoS 1/2 PUBLISH outside the new limits | Not resent and removed from the session; settled with `indeterminate: may have been delivered; not resent because the new connection's limits do not allow it ()`. A QoS 2 packet identifier abandoned this way is not reused for new messages until a connection reports Session Present = 0 | +| Session Present = 1 | QoS 2 awaiting PUBCOMP | PUBREL resent | +| Session Present = 0 (server lost the session) | QoS 1 PUBLISH within the new limits | Sent again as a new message (DUP = 0), in the original order; its promise or callback settles normally when acknowledged | +| Session Present = 0 | QoS 1 PUBLISH outside the new limits | Not sent; settled with `indeterminate: session lost; may have been delivered; not resent because the new connection's limits do not allow it ()` | +| Session Present = 0 | QoS 2 PUBLISH or PUBREL | Not sent; settled with `indeterminate: session lost; may have been delivered` | + +`publishWithOptions` promises reject with `Publish not acknowledged: ` followed by that text, and `publishQos1`/`publishQos2` callbacks receive the text as their only argument. Messages whose promise or callback was already settled by `disconnect()` are handled the same way without a further notification. + +A CONNACK with Session Present = 1 is rejected (the client sends DISCONNECT 0x82 on v5.0 and closes the connection) unless the client holds session state or `resumeExistingSession` is set. + ### PublishOptions API The JavaScript class is exported as `PublishOptions` (Rust type: `WasmPublishOptions`). diff --git a/crates/mqtt5-conformance/CONFORMANCE_DIARY.md b/crates/mqtt5-conformance/CONFORMANCE_DIARY.md index 94b649c4..edbe93ec 100644 --- a/crates/mqtt5-conformance/CONFORMANCE_DIARY.md +++ b/crates/mqtt5-conformance/CONFORMANCE_DIARY.md @@ -38,11 +38,21 @@ ## Diary Entries +### Quorum review of the client fixes, and a TLA+-verified outcome model for the offline queue (2026-09-23) + +**Trigger**: a five-reviewer quorum review of PR #164 before merge. Most findings came with a failing test. The worst was a regression in the wasm client: a v3.1.1 persistent session could never reconnect, because the session-lifetime check used Session Expiry (always 0 in 3.1.1) and the new strict `[MQTT-3.2.2-4]` check then rejected the broker's Session Present=1 forever. Other findings: QUIC teardown left the connection open after the client's own DISCONNECT; a publish waiting across a reconnect was sent under the old server's limits; and the offline queue dropped messages silently after `publish()` had returned success. All were fixed in the same PR. + +**The offline-queue question was settled by three independent TLA+ models, not by argument.** All three found the same defects in the current behaviour: silent loss at flush, silent loss on Session Present=0, and replay resending packets the new CONNACK forbids. All three arrived at the same design: per-publish outcomes (Delivered / Rejected / Indeterminate) that are never keyed by packet id, checks at enqueue, reject-and-continue at flush, re-checks on replay, and quarantine of abandoned QoS 2 ids. Holding a message fails liveness, skipping ahead breaks ordering, and transforming it silently loses the caller's intent. One model showed that reporting failures by packet id misattributes them after the id is reused (ABA). The user chose downgrade-and-report for Maximum QoS, and requeueing QoS 1 when the server loses the session. A Clean Start=1 connect discards unacked outbound state, as `[MQTT-3.1.2-4]` requires. The consolidated spec is in `specs/tla/offline-queue/`. New tests are in `crates/mqtt5/tests/conf_client_offline_queue.rs`: 11 fake-broker tests, all failing on the prior tree. + +**Second quorum round on the outcome code**: three re-reviewers checked the implementation against the models. The protocol behaviour held. Two defects were in how outcomes are settled: a connection loss never closed the send quota, so a parked flush task kept handles alive forever, and a reader aborted between releasing session state and settling the outcome left a delivered publish unsettled. A live publish returned `Err` on connection loss although its message stayed in the session and was resent. It now returns a handle like a queued one. A QoS 2 message that received PUBREC Success before the session was lost is reported delivered. Each fix has a test that fails on the previous code. + +**Tooling lesson**: tla-mcp 0.9.4 passed a `~>` negative control vacuously and ignored missing fairness. Liveness results are trusted only as `[]<>` properties with negative controls that fail, backed by an ENABLED-based progress invariant. + ### Client-side audit: the suite only ever tested brokers, and our own clients failed ~40 MUSTs (2026-09-23) **Trigger**: checking a third-party client's claim of full MQTT v5 conformance. This suite's SUT is always a broker, so it could not answer the question. The 149 statements in `conformance.toml` with `applies_to = "Client"` or `"Both"` were audited instead with raw-byte fake-broker tests that drive the real client and record what it puts on the wire. After the third-party client had been tested, the same audit ran against our own `MqttClient` and the `mqtt5-wasm` client. -**Result**: our native client failed about 40 client MUST statements. The third-party client failed 7. Resend on session resume did not exist. Packet identifiers were reused while in flight. There was no topic or filter validation. Topic Alias Maximum and Retain Available were never enforced. The offline queue bypassed flow control. Protocol errors left the socket half-open. WebSocket reads assumed one packet per frame. The wasm client had most of the same defects and also never sent PUBACK. All of them are fixed in mqtt5 0.41.0 and mqtt5-wasm 1.5.0. +**Result**: our native client failed about 40 client MUST statements. The third-party client failed 7. Resend on session resume did not exist. Packet identifiers were reused while in flight. There was no topic or filter validation. Topic Alias Maximum and Retain Available were never enforced. The offline queue bypassed flow control. Protocol errors left the socket half-open. WebSocket reads assumed one packet per frame. The wasm client had most of the same defects and also never sent PUBACK. All of them are fixed in mqtt5 0.41.0 and mqtt5-wasm 2.0.0. **Where the tests live**: `crates/mqtt5/tests/conf_client_{a,b,c,d}.rs` (native) and `crates/mqtt5-wasm/tests/conformance_client.rs` (wasm, MessagePort fake broker under Node). They are not yet registered in this crate's manifest or runner. A client-side SUT mode for this suite is the natural next step. diff --git a/crates/mqtt5-conformance/src/test_client/inprocess.rs b/crates/mqtt5-conformance/src/test_client/inprocess.rs index 80583aff..5915d723 100644 --- a/crates/mqtt5-conformance/src/test_client/inprocess.rs +++ b/crates/mqtt5-conformance/src/test_client/inprocess.rs @@ -62,18 +62,25 @@ impl InProcessTestClient { /// Publishes with the given [`PublishOptions`]. /// /// # Errors - /// Returns an error if the broker rejects the publish or the client - /// is disconnected. + /// Returns an error if the broker rejects the publish, the client is + /// disconnected, or the publish was not acknowledged before the connection + /// ended or the acknowledgement wait elapsed. pub async fn publish_with_options( &self, topic: &str, payload: &[u8], options: PublishOptions, ) -> Result<(), TestClientError> { - self.client + match self + .client .publish_with_options(topic, payload.to_vec(), options) - .await?; - Ok(()) + .await? + { + mqtt5::PublishResult::Sent(_) => Ok(()), + mqtt5::PublishResult::Queued(_) => { + Err(TestClientError::Timeout("publish acknowledgement")) + } + } } /// Subscribes to `filter` and returns a [`Subscription`] handle. diff --git a/crates/mqtt5-wasm/Cargo.toml b/crates/mqtt5-wasm/Cargo.toml index 1e5c935b..bfc28f82 100644 --- a/crates/mqtt5-wasm/Cargo.toml +++ b/crates/mqtt5-wasm/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "mqtt5-wasm" -version = "1.5.0" +version = "2.0.0" edition.workspace = true rust-version.workspace = true authors.workspace = true diff --git a/crates/mqtt5-wasm/README.md b/crates/mqtt5-wasm/README.md index 36edca0e..d51aed92 100644 --- a/crates/mqtt5-wasm/README.md +++ b/crates/mqtt5-wasm/README.md @@ -29,7 +29,7 @@ npm install mqtt5-wasm ```toml [dependencies] -mqtt5-wasm = "1.3" +mqtt5-wasm = "2.0" ``` Build with wasm-bindgen: diff --git a/crates/mqtt5-wasm/src/client/callbacks.rs b/crates/mqtt5-wasm/src/client/callbacks.rs index 97aea9e5..03fccaf6 100644 --- a/crates/mqtt5-wasm/src/client/callbacks.rs +++ b/crates/mqtt5-wasm/src/client/callbacks.rs @@ -11,19 +11,42 @@ use super::reconnect::spawn_reconnection_task; use super::state::ClientState; const SESSION_DISCARDED_REASON: u8 = 0x80; +const DISCONNECTED_WITH_SESSION: &str = + "disconnected; message remains in session and will be resent on resume by this client instance"; -fn reject_callbacks(callbacks: Vec) { - let error_val = JsValue::from_f64(f64::from(SESSION_DISCARDED_REASON)); +pub const SESSION_DISCARDED_BY_CLEAN_START: &str = + "indeterminate: session discarded by clean start; may have been delivered"; +const SESSION_ENDED_WITH_CONNECTION: &str = + "indeterminate: session ended with the connection; may have been delivered"; + +fn settle_callbacks(callbacks: Vec, message: &str) { + let message = JsValue::from_str(message); for callback in callbacks { - if let Err(e) = callback.call1(&JsValue::NULL, &error_val) { + if let Err(e) = callback.call1(&JsValue::NULL, &message) { tracing::warn!(error = ?e, "acknowledgement callback failed"); } } } -pub fn discard_session(state: &Rc>) { - let callbacks = state.borrow_mut().discard_session(); - reject_callbacks(callbacks); +pub fn settle_acks_on_disconnect(state: &Rc>) { + let callbacks = state.borrow_mut().take_ack_callbacks(); + settle_callbacks(callbacks, DISCONNECTED_WITH_SESSION); +} + +pub fn discard_session(state: &Rc>, message: &str) { + let (callbacks, abandoned) = { + let mut state_mut = state.borrow_mut(); + let abandoned = state_mut.outbound.len(); + (state_mut.discard_session(), abandoned) + }; + if abandoned > 0 { + tracing::warn!( + abandoned, + reason = message, + "outbound session state discarded" + ); + } + settle_callbacks(callbacks, message); } pub fn close_network_connection(state: &Rc>) { @@ -49,7 +72,7 @@ pub fn end_connection_state(state: &Rc>) { .drain() .filter_map(|(_, callback)| callback) .collect(); - (subacks, state_mut.session_expiry_interval == 0) + (subacks, !state_mut.session_outlives_connection()) }; for resolve in subacks { let codes = js_sys::Array::new(); @@ -59,7 +82,7 @@ pub fn end_connection_state(state: &Rc>) { } } if session_ends { - discard_session(state); + discard_session(state, SESSION_ENDED_WITH_CONNECTION); } wake_quota_waiters(state); } diff --git a/crates/mqtt5-wasm/src/client/connection.rs b/crates/mqtt5-wasm/src/client/connection.rs index b016d161..8bd6d670 100644 --- a/crates/mqtt5-wasm/src/client/connection.rs +++ b/crates/mqtt5-wasm/src/client/connection.rs @@ -12,11 +12,11 @@ use wasm_bindgen::JsValue; use crate::transport::{WasmReader, WasmTransportType}; -use super::callbacks::discard_session; +use super::callbacks::{discard_session, SESSION_DISCARDED_BY_CLEAN_START}; use super::handlers::handle_auth; use super::keepalive::spawn_keepalive_task; use super::packet::{encode_packet, write_packet}; -use super::qos::{resend_session, spawn_qos2_cleanup_task}; +use super::qos::{resume_session, spawn_qos2_cleanup_task}; use super::reader::{read_connect_response, spawn_packet_reader}; use super::state::{ClientState, SessionState, StoredConnectOptions}; @@ -66,9 +66,12 @@ pub async fn establish( .map_err(|e| ConnectFailure::Failed(format!("Transport connection failed: {e}")))?; if connect.clean_start { - discard_session(state); + state.borrow_mut().quarantined.clear(); + discard_session(state, SESSION_DISCARDED_BY_CLEAN_START); } - state.borrow_mut().apply_connect_options(options); + state + .borrow_mut() + .apply_connect_options(options, connect.clean_start); let mut buf = BytesMut::new(); encode_packet(&Packet::Connect(Box::new(connect.clone())), &mut buf) @@ -177,10 +180,6 @@ fn accept_connack( )); } - if !connack.session_present { - discard_session(state); - } - { let mut state_mut = state.borrow_mut(); state_mut.apply_connack(&connack); @@ -189,9 +188,7 @@ fn accept_connack( state_mut.connection_generation = state_mut.connection_generation.wrapping_add(1); } - if connack.session_present { - resend_session(state); - } + resume_session(state, connack.session_present); spawn_packet_reader(Rc::clone(state), reader); spawn_keepalive_task(Rc::clone(state)); diff --git a/crates/mqtt5-wasm/src/client/handlers.rs b/crates/mqtt5-wasm/src/client/handlers.rs index fff07342..47afd5ae 100644 --- a/crates/mqtt5-wasm/src/client/handlers.rs +++ b/crates/mqtt5-wasm/src/client/handlers.rs @@ -357,12 +357,14 @@ fn complete_flight( state: &Rc>, packet_id: u16, reason_code: ReasonCode, + qos: QoS, callback: Option, ) { release_quota(state); if let Some(callback) = callback { let reason_code_js = JsValue::from_f64(f64::from(u8::from(reason_code))); - if let Err(e) = callback.call1(&JsValue::NULL, &reason_code_js) { + let qos_js = JsValue::from_f64(f64::from(qos as u8)); + if let Err(e) = callback.call2(&JsValue::NULL, &reason_code_js, &qos_js) { tracing::warn!(error = ?e, packet_id, "publish acknowledgement callback failed"); } } @@ -382,7 +384,7 @@ fn handle_puback(state: &Rc>, packet_id: u16, reason_code: state_mut.outbound.remove(&packet_id); state_mut.pending_pubacks.remove(&packet_id) }; - complete_flight(state, packet_id, reason_code, callback); + complete_flight(state, packet_id, reason_code, QoS::AtLeastOnce, callback); } enum PubRecOutcome { @@ -414,7 +416,9 @@ fn handle_pubrec(state: &Rc>, packet_id: u16, reason_code: }; match outcome { PubRecOutcome::Release => send_ack(state, &Packet::PubRel(PubRelPacket::new(packet_id))), - PubRecOutcome::Failed(callback) => complete_flight(state, packet_id, reason_code, callback), + PubRecOutcome::Failed(callback) => { + complete_flight(state, packet_id, reason_code, QoS::ExactlyOnce, callback); + } PubRecOutcome::Unknown => send_ack( state, &Packet::PubRel(PubRelPacket::new_with_reason( @@ -442,7 +446,7 @@ fn handle_pubcomp(state: &Rc>, packet_id: u16, reason_code: .remove(&packet_id) .map(|(callback, _)| callback) }; - complete_flight(state, packet_id, reason_code, callback); + complete_flight(state, packet_id, reason_code, QoS::ExactlyOnce, callback); } fn handle_pubrel(state: &Rc>, packet_id: u16) { diff --git a/crates/mqtt5-wasm/src/client/mod.rs b/crates/mqtt5-wasm/src/client/mod.rs index 8d2de9d7..0cbe39a2 100644 --- a/crates/mqtt5-wasm/src/client/mod.rs +++ b/crates/mqtt5-wasm/src/client/mod.rs @@ -29,9 +29,12 @@ use wasm_bindgen::prelude::*; use wasm_bindgen_futures::JsFuture; use web_sys::MessagePort; -use callbacks::{close_network_connection, end_connection_state, trigger_disconnect_callback}; +use callbacks::{ + close_network_connection, end_connection_state, settle_acks_on_disconnect, + trigger_disconnect_callback, +}; use connection::establish; -use outbound::{check_publish, check_subscribe, check_unsubscribe}; +use outbound::{check_publish, check_subscribe, check_unsubscribe, downgrade_to_server_maximum}; use packet::write_packet; use qos::{abandon_flight, await_ack_promises, create_ack_promises, reserve_flight}; use state::{ClientState, StoredConnectOptions}; @@ -215,7 +218,7 @@ impl WasmMqttClient { topic: &str, payload: &[u8], options: &WasmPublishOptions, - ) -> Result<(), JsValue> { + ) -> Result { let qos = options.to_qos(); #[cfg(feature = "codec")] @@ -263,11 +266,13 @@ impl WasmMqttClient { stream_id: None, }; - let packet_id = self + let (packet_id, qos_used) = self .dispatch_publish(publish_packet, AckSink::Promise) .await?; - let (puback_promise, pubcomp_promise) = create_ack_promises(&self.state, qos, packet_id); - await_ack_promises(puback_promise, pubcomp_promise).await + let (puback_promise, pubcomp_promise) = + create_ack_promises(&self.state, qos_used, packet_id); + await_ack_promises(puback_promise, pubcomp_promise).await?; + Ok(qos_used as u8) } /// # Errors @@ -419,6 +424,7 @@ impl WasmMqttClient { } close_network_connection(&self.state); end_connection_state(&self.state); + settle_acks_on_disconnect(&self.state); trigger_disconnect_callback(&self.state); Ok(()) } @@ -563,41 +569,49 @@ impl WasmMqttClient { protocol_version, stream_id: None, }; - self.dispatch_publish(publish_packet, AckSink::Callback(callback)) - .await? - .ok_or_else(|| js_error("QoS 0 publish has no packet identifier")) + let (packet_id, _) = self + .dispatch_publish(publish_packet, AckSink::Callback(callback)) + .await?; + Ok(packet_id.unwrap_or(0)) } async fn dispatch_publish( &self, mut publish: PublishPacket, sink: AckSink, - ) -> Result, JsValue> { + ) -> Result<(Option, QoS), JsValue> { self.ensure_connected().await?; - check_publish(&self.state.borrow(), &publish).map_err(js_error)?; + let generation = { + let state = self.state.borrow(); + publish.protocol_version = state.protocol_version; + downgrade_to_server_maximum(&state.server, &mut publish); + check_publish(&state, &publish).map_err(js_error)?; + state.connection_generation + }; + let qos = publish.qos; - let packet_id = if publish.qos == QoS::AtMostOnce { + let packet_id = if qos == QoS::AtMostOnce { None } else { - let packet_id = reserve_flight(&self.state, &publish).await?; - if let Err(e) = check_publish(&self.state.borrow(), &publish) { - abandon_flight(&self.state, packet_id); - return Err(js_error(e)); - } - Some(packet_id) + Some(reserve_flight(&self.state, &publish, generation).await?) }; publish.packet_id = packet_id; - if let (Some(packet_id), AckSink::Callback(callback)) = (packet_id, sink) { - let mut state = self.state.borrow_mut(); - if publish.qos == QoS::ExactlyOnce { - state - .pending_pubcomps - .insert(packet_id, (callback, js_sys::Date::now())); - } else { - state.pending_pubacks.insert(packet_id, callback); + let immediate_callback = match (packet_id, sink) { + (Some(packet_id), AckSink::Callback(callback)) => { + let mut state = self.state.borrow_mut(); + if qos == QoS::ExactlyOnce { + state + .pending_pubcomps + .insert(packet_id, (callback, js_sys::Date::now())); + } else { + state.pending_pubacks.insert(packet_id, callback); + } + None } - } + (None, AckSink::Callback(callback)) => Some(callback), + (_, AckSink::Promise) => None, + }; let alias_mapping = publish .topic_alias() @@ -627,7 +641,15 @@ impl WasmMqttClient { } } - Ok(packet_id) + if let Some(callback) = immediate_callback { + let reason_code = JsValue::from_f64(0.0); + let qos_used = JsValue::from_f64(f64::from(qos as u8)); + if let Err(e) = callback.call2(&JsValue::NULL, &reason_code, &qos_used) { + tracing::warn!(error = ?e, "publish completion callback failed"); + } + } + + Ok((packet_id, qos)) } async fn send_subscribe( @@ -747,12 +769,14 @@ impl WasmMqttClient { topic: &str, payload: &[u8], qos: QoS, - ) -> Result<(), JsValue> { + ) -> Result { let publish_packet = PublishPacket::new(topic.to_string(), payload.to_vec(), qos); - let packet_id = self + let (packet_id, qos_used) = self .dispatch_publish(publish_packet, AckSink::Promise) .await?; - let (puback_promise, pubcomp_promise) = create_ack_promises(&self.state, qos, packet_id); - await_ack_promises(puback_promise, pubcomp_promise).await + let (puback_promise, pubcomp_promise) = + create_ack_promises(&self.state, qos_used, packet_id); + await_ack_promises(puback_promise, pubcomp_promise).await?; + Ok(qos_used) } } diff --git a/crates/mqtt5-wasm/src/client/outbound.rs b/crates/mqtt5-wasm/src/client/outbound.rs index 9f47d694..50453bca 100644 --- a/crates/mqtt5-wasm/src/client/outbound.rs +++ b/crates/mqtt5-wasm/src/client/outbound.rs @@ -10,9 +10,47 @@ use mqtt5_protocol::validation::{ use mqtt5_protocol::QoS; use super::packet::encode_checked; -use super::state::ClientState; +use super::state::{ClientState, ServerLimits}; pub fn check_publish(state: &ClientState, publish: &PublishPacket) -> Result<(), String> { + if state.protocol_version == 5 { + check_v5_publish_properties(state, publish)?; + } else { + validate_topic_name(&publish.topic_name).map_err(|e| e.to_string())?; + } + check_server_limits(&state.server, publish) +} + +pub fn check_server_limits(server: &ServerLimits, publish: &PublishPacket) -> Result<(), String> { + if publish.qos as u8 > server.maximum_qos as u8 { + return Err(format!( + "QoS {} exceeds the server Maximum QoS {}", + publish.qos as u8, server.maximum_qos as u8 + )); + } + if publish.retain && !server.retain_available { + return Err("The server does not support retained messages".to_string()); + } + let mut sized = publish.clone(); + if sized.qos != QoS::AtMostOnce { + sized.packet_id = Some(sized.packet_id.unwrap_or(u16::MAX)); + } + encode_checked(&Packet::Publish(sized), server.maximum_packet_size)?; + Ok(()) +} + +pub fn downgrade_to_server_maximum(server: &ServerLimits, publish: &mut PublishPacket) { + if publish.qos as u8 > server.maximum_qos as u8 { + tracing::warn!( + requested = publish.qos as u8, + maximum = server.maximum_qos as u8, + "requested QoS exceeds the server Maximum QoS; publishing at the maximum" + ); + publish.qos = server.maximum_qos; + } +} + +fn check_v5_publish_properties(state: &ClientState, publish: &PublishPacket) -> Result<(), String> { let alias = publish.topic_alias(); match (publish.topic_name.is_empty(), alias) { (true, None) => return Err("A zero-length Topic Name requires a Topic Alias".to_string()), @@ -31,20 +69,6 @@ pub fn check_publish(state: &ClientState, publish: &PublishPacket) -> Result<(), { return Err("A client PUBLISH must not contain a Subscription Identifier".to_string()); } - if publish.qos as u8 > state.server.maximum_qos as u8 { - return Err(format!( - "QoS {} exceeds the server Maximum QoS {}", - publish.qos as u8, state.server.maximum_qos as u8 - )); - } - if publish.retain && !state.server.retain_available { - return Err("The server does not support retained messages".to_string()); - } - let mut sized = publish.clone(); - if sized.qos != QoS::AtMostOnce { - sized.packet_id = Some(sized.packet_id.unwrap_or(u16::MAX)); - } - encode_checked(&Packet::Publish(sized), state.server.maximum_packet_size)?; Ok(()) } @@ -118,3 +142,44 @@ pub fn check_unsubscribe(state: &ClientState, packet: &UnsubscribePacket) -> Res )?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + use wasm_bindgen_test::wasm_bindgen_test; + + fn state_for(protocol_version: u8) -> ClientState { + let mut state = ClientState::new("outbound".to_string()); + state.protocol_version = protocol_version; + state + } + + fn publish_with_v5_only_properties(protocol_version: u8, topic: &str) -> PublishPacket { + let mut publish = PublishPacket::new(topic, b"x".to_vec(), QoS::AtLeastOnce); + publish.protocol_version = protocol_version; + publish.properties.set_topic_alias(3); + publish.properties.set_response_topic("reply/+".to_string()); + publish.properties.set_subscription_identifier(7); + publish + } + + #[wasm_bindgen_test] + fn v311_publish_skips_v5_only_property_checks() { + let publish = publish_with_v5_only_properties(4, "a/b"); + assert_eq!(check_publish(&state_for(4), &publish), Ok(())); + } + + #[wasm_bindgen_test] + fn v311_publish_still_validates_topic_name() { + let publish = publish_with_v5_only_properties(4, "a/+"); + assert!(check_publish(&state_for(4), &publish).is_err()); + let publish = publish_with_v5_only_properties(4, ""); + assert!(check_publish(&state_for(4), &publish).is_err()); + } + + #[wasm_bindgen_test] + fn v5_publish_applies_v5_only_property_checks() { + let publish = publish_with_v5_only_properties(5, "a/b"); + assert!(check_publish(&state_for(5), &publish).is_err()); + } +} diff --git a/crates/mqtt5-wasm/src/client/qos.rs b/crates/mqtt5-wasm/src/client/qos.rs index 07a354f1..a7318a0b 100644 --- a/crates/mqtt5-wasm/src/client/qos.rs +++ b/crates/mqtt5-wasm/src/client/qos.rs @@ -7,11 +7,14 @@ use std::rc::Rc; use wasm_bindgen::prelude::*; use wasm_bindgen_futures::JsFuture; +use super::outbound::check_server_limits; use super::packet::write_packet; use super::sleep_ms; use super::state::ClientState; const QOS2_CALLBACK_TIMEOUT_MS: f64 = 10_000.0; +const INDETERMINATE: &str = "indeterminate: may have been delivered"; +const SESSION_LOST: &str = "indeterminate: session lost; may have been delivered"; pub fn create_ack_promises( state: &Rc>, @@ -87,6 +90,7 @@ fn stored_copy(state: &ClientState, publish: &PublishPacket) -> PublishPacket { pub async fn reserve_flight( state: &Rc>, publish: &PublishPacket, + generation: u32, ) -> Result { loop { { @@ -94,6 +98,11 @@ pub async fn reserve_flight( if !state_mut.connected { return Err(JsValue::from_str("Not connected")); } + if state_mut.connection_generation != generation { + return Err(JsValue::from_str( + "Connection changed before the message was sent", + )); + } if state_mut.send_quota > 0 { let packet_id = state_mut .allocate_packet_id() @@ -123,14 +132,19 @@ pub fn abandon_flight(state: &Rc>, packet_id: u16) { pub fn release_quota(state: &Rc>) { let (resend, waiter) = { let mut state_mut = state.borrow_mut(); + if state_mut.quota_debt > 0 { + state_mut.quota_debt -= 1; + return; + } state_mut.send_quota = state_mut .send_quota .saturating_add(1) .min(state_mut.server.receive_maximum); let mut resend = None; while let Some(packet_id) = state_mut.pending_resends.pop_front() { - if let Some(flight) = state_mut.outbound.get(&packet_id) { - resend = Some(flight.publish.clone().with_dup(true)); + if let Some(flight) = state_mut.outbound.get_mut(&packet_id) { + resend = Some(flight.publish.clone().with_dup(flight.transmitted)); + flight.transmitted = true; break; } } @@ -163,29 +177,103 @@ pub fn wake_quota_waiters(state: &Rc>) { } } -pub fn resend_session(state: &Rc>) { +pub fn resume_session(state: &Rc>, session_present: bool) { + let settlements = { + let mut state_mut = state.borrow_mut(); + if !session_present { + state_mut.awaiting_pubrel.clear(); + state_mut.quarantined.clear(); + } + let mut settlements = Vec::new(); + for packet_id in state_mut.flights_in_order() { + let Some(verdict) = session_verdict(&mut state_mut, packet_id, session_present) else { + continue; + }; + state_mut.outbound.remove(&packet_id); + if verdict.quarantine { + state_mut.quarantined.insert(packet_id); + } + tracing::warn!(packet_id, reason = %verdict.message, "unacknowledged PUBLISH abandoned"); + if let Some(callback) = state_mut.take_ack_callback(packet_id) { + settlements.push((callback, verdict.message)); + } + } + settlements + }; + resend_session(state); + for (callback, message) in settlements { + if let Err(e) = callback.call1(&JsValue::NULL, &JsValue::from_str(&message)) { + tracing::warn!(error = ?e, "acknowledgement callback failed"); + } + } +} + +struct Abandon { + message: String, + quarantine: bool, +} + +fn session_verdict( + state: &mut ClientState, + packet_id: u16, + session_present: bool, +) -> Option { + let flight = state.outbound.get_mut(&packet_id)?; + let not_resent = |reason: String, prefix: &str| { + format!( + "{prefix}; not resent because the new connection's limits do not allow it ({reason})" + ) + }; + match (session_present, flight.publish.qos, flight.released) { + (true, _, true) => None, + (true, qos, false) => check_server_limits(&state.server, &flight.publish) + .err() + .map(|reason| Abandon { + message: not_resent(reason, INDETERMINATE), + quarantine: qos == QoS::ExactlyOnce, + }), + (false, QoS::ExactlyOnce, _) => Some(Abandon { + message: SESSION_LOST.to_string(), + quarantine: false, + }), + (false, _, _) => { + flight.transmitted = false; + check_server_limits(&state.server, &flight.publish) + .err() + .map(|reason| Abandon { + message: not_resent(reason, SESSION_LOST), + quarantine: false, + }) + } + } +} + +fn resend_session(state: &Rc>) { let packets = { let mut state_mut = state.borrow_mut(); - let mut order: Vec<(u64, u16)> = state_mut - .outbound - .iter() - .map(|(packet_id, flight)| (flight.sequence, *packet_id)) - .collect(); - order.sort_unstable(); - let mut packets = Vec::with_capacity(order.len()); - for (_, packet_id) in order { - let Some(flight) = state_mut.outbound.get(&packet_id) else { + let mut packets = Vec::with_capacity(state_mut.outbound.len()); + for packet_id in state_mut.flights_in_order() { + let mut quota = state_mut.send_quota; + let Some(flight) = state_mut.outbound.get_mut(&packet_id) else { continue; }; if flight.released { packets.push(Packet::PubRel(PubRelPacket::new(packet_id))); - state_mut.send_quota = state_mut.send_quota.saturating_sub(1); - } else if state_mut.send_quota > 0 { - packets.push(Packet::Publish(flight.publish.clone().with_dup(true))); - state_mut.send_quota -= 1; + if quota > 0 { + quota -= 1; + } else { + state_mut.quota_debt = state_mut.quota_debt.saturating_add(1); + } + } else if quota > 0 { + packets.push(Packet::Publish( + flight.publish.clone().with_dup(flight.transmitted), + )); + flight.transmitted = true; + quota -= 1; } else { state_mut.pending_resends.push_back(packet_id); } + state_mut.send_quota = quota; } packets }; diff --git a/crates/mqtt5-wasm/src/client/reconnect.rs b/crates/mqtt5-wasm/src/client/reconnect.rs index 19cc7d6f..d9168843 100644 --- a/crates/mqtt5-wasm/src/client/reconnect.rs +++ b/crates/mqtt5-wasm/src/client/reconnect.rs @@ -11,7 +11,7 @@ use super::callbacks::{trigger_reconnect_failed_callback, trigger_reconnecting_c use super::connection::establish; use super::connectivity::{is_browser_online, wait_for_online}; use super::sleep_ms; -use super::state::{ClientState, StoredConnectOptions}; +use super::state::{ClientState, SessionState, StoredConnectOptions}; pub fn spawn_reconnection_task(state: Rc>) { spawn_local(async move { @@ -118,7 +118,13 @@ async fn attempt_reconnect( let transport = WasmTransportType::WebSocket( crate::transport::websocket::WasmWebSocketTransport::new(url), ); - let client_id = state.borrow().client_id.clone(); + let (client_id, clean_start) = { + let state_ref = state.borrow(); + ( + state_ref.client_id.clone(), + state_ref.session == SessionState::Absent && !options.resume_existing_session, + ) + }; let properties = if options.protocol_version == 5 { build_properties_from_stored(options) } else { @@ -127,7 +133,7 @@ async fn attempt_reconnect( let connect_packet = ConnectPacket { protocol_version: options.protocol_version, - clean_start: false, + clean_start, keep_alive: options.keep_alive, client_id, username: options.username.clone(), diff --git a/crates/mqtt5-wasm/src/client/state.rs b/crates/mqtt5-wasm/src/client/state.rs index 3408ece7..5b517e05 100644 --- a/crates/mqtt5-wasm/src/client/state.rs +++ b/crates/mqtt5-wasm/src/client/state.rs @@ -15,6 +15,7 @@ use std::rc::Rc; use crate::codec::WasmCodecRegistry; const DEFAULT_RECEIVE_MAXIMUM: u16 = u16::MAX; +const SESSION_NEVER_EXPIRES: u32 = u32::MAX; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct ServerLimits { @@ -118,6 +119,7 @@ pub struct OutboundFlight { pub sequence: u64, pub publish: PublishPacket, pub released: bool, + pub transmitted: bool, } pub struct ClientState { @@ -135,9 +137,11 @@ pub struct ClientState { pub outbound: HashMap, pub next_flight_sequence: u64, pub send_quota: u16, + pub quota_debt: u16, pub quota_waiters: VecDeque, pub pending_resends: VecDeque, pub awaiting_pubrel: HashSet, + pub quarantined: HashSet, pub server: ServerLimits, pub client_limits: ClientLimits, pub outbound_aliases: TopicAliasManager, @@ -186,9 +190,11 @@ impl ClientState { outbound: HashMap::new(), next_flight_sequence: 0, send_quota: DEFAULT_RECEIVE_MAXIMUM, + quota_debt: 0, quota_waiters: VecDeque::new(), pending_resends: VecDeque::new(), awaiting_pubrel: HashSet::new(), + quarantined: HashSet::new(), server: ServerLimits::default(), client_limits: ClientLimits::default(), outbound_aliases: TopicAliasManager::new(0), @@ -225,6 +231,7 @@ impl ClientState { self.outbound.contains_key(&packet_id) || self.pending_subacks.contains_key(&packet_id) || self.pending_unsubacks.contains(&packet_id) + || self.quarantined.contains(&packet_id) } pub fn allocate_packet_id(&self) -> Option { @@ -243,15 +250,39 @@ impl ClientState { sequence, publish, released: false, + transmitted: true, }, ); } + pub fn flights_in_order(&self) -> Vec { + let mut order: Vec<(u64, u16)> = self + .outbound + .iter() + .map(|(packet_id, flight)| (flight.sequence, *packet_id)) + .collect(); + order.sort_unstable(); + order.into_iter().map(|(_, packet_id)| packet_id).collect() + } + + pub fn take_ack_callback(&mut self, packet_id: u16) -> Option { + self.pending_pubacks.remove(&packet_id).or_else(|| { + self.pending_pubcomps + .remove(&packet_id) + .map(|(callback, _)| callback) + }) + } + pub fn discard_session(&mut self) -> Vec { self.outbound.clear(); self.pending_resends.clear(); self.awaiting_pubrel.clear(); + self.quota_debt = 0; self.session = SessionState::Absent; + self.take_ack_callbacks() + } + + pub fn take_ack_callbacks(&mut self) -> Vec { self.pending_pubacks .drain() .map(|(_, callback)| callback) @@ -263,10 +294,18 @@ impl ClientState { .collect() } - pub fn apply_connect_options(&mut self, options: &StoredConnectOptions) { + pub fn session_outlives_connection(&self) -> bool { + self.session_expiry_interval > 0 + } + + pub fn apply_connect_options(&mut self, options: &StoredConnectOptions, clean_start: bool) { self.keep_alive = options.keep_alive; self.protocol_version = options.protocol_version; - self.session_expiry_interval = options.session_expiry_interval.unwrap_or(0); + self.session_expiry_interval = match (options.protocol_version, clean_start) { + (5, _) => options.session_expiry_interval.unwrap_or(0), + (_, true) => 0, + (_, false) => SESSION_NEVER_EXPIRES, + }; self.client_limits = ClientLimits::from(options); self.auth_method.clone_from(&options.authentication_method); #[cfg(feature = "codec")] @@ -278,6 +317,10 @@ impl ClientState { pub fn apply_connack(&mut self, connack: &ConnAckPacket) { self.server = ServerLimits::from_connack(connack); self.send_quota = self.server.receive_maximum; + self.quota_debt = 0; + if let Some(interval) = connack.properties.get_session_expiry_interval() { + self.session_expiry_interval = interval; + } self.outbound_aliases = TopicAliasManager::new(self.server.topic_alias_maximum); self.inbound_aliases = TopicAliasManager::new(self.client_limits.topic_alias_maximum); self.pending_resends.clear(); diff --git a/crates/mqtt5-wasm/src/client_handler/connect.rs b/crates/mqtt5-wasm/src/client_handler/connect.rs index c3fc14ff..87750a9c 100644 --- a/crates/mqtt5-wasm/src/client_handler/connect.rs +++ b/crates/mqtt5-wasm/src/client_handler/connect.rs @@ -1,4 +1,5 @@ use mqtt5::broker::auth::{EnhancedAuthResult, EnhancedAuthStatus}; +use mqtt5::broker::router::SubscriptionRequest; use mqtt5::broker::storage::{ClientSession, StorageBackend}; use mqtt5_protocol::error::{MqttError, Result}; use mqtt5_protocol::packet::auth::AuthPacket; @@ -227,16 +228,16 @@ impl WasmClientHandler { } self.router .subscribe( - client_id.clone(), - topic_filter.clone(), - stored.qos, - stored.subscription_id, - stored.no_local, - stored.retain_as_published, - stored.retain_handling, - ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), - stored.change_only, - stored.flow_id, + SubscriptionRequest::new(client_id.clone(), topic_filter.clone(), stored.qos) + .with_subscription_id(stored.subscription_id) + .with_no_local(stored.no_local) + .with_retain_as_published(stored.retain_as_published) + .with_retain_handling(stored.retain_handling) + .with_protocol_version( + ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), + ) + .with_change_only(stored.change_only) + .with_flow_id(stored.flow_id), ) .await?; } diff --git a/crates/mqtt5-wasm/src/client_handler/subscribe.rs b/crates/mqtt5-wasm/src/client_handler/subscribe.rs index c481565d..0ee0c98d 100644 --- a/crates/mqtt5-wasm/src/client_handler/subscribe.rs +++ b/crates/mqtt5-wasm/src/client_handler/subscribe.rs @@ -1,3 +1,4 @@ +use mqtt5::broker::router::SubscriptionRequest; use mqtt5::broker::storage::{StorageBackend, StoredSubscription}; use mqtt5_protocol::error::{MqttError, Result}; use mqtt5_protocol::packet::disconnect::DisconnectPacket; @@ -91,16 +92,15 @@ impl WasmClientHandler { self.router .subscribe( - client_id.clone(), - filter.filter.clone(), - granted_qos, - subscription_id, - filter.options.no_local, - filter.options.retain_as_published, - filter.options.retain_handling as u8, - ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), - change_only, - None, + SubscriptionRequest::new(client_id.clone(), filter.filter.clone(), granted_qos) + .with_subscription_id(subscription_id) + .with_no_local(filter.options.no_local) + .with_retain_as_published(filter.options.retain_as_published) + .with_retain_handling(filter.options.retain_handling as u8) + .with_protocol_version( + ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), + ) + .with_change_only(change_only), ) .await?; diff --git a/crates/mqtt5-wasm/tests/client_broker.rs b/crates/mqtt5-wasm/tests/client_broker.rs new file mode 100644 index 00000000..eb77835a --- /dev/null +++ b/crates/mqtt5-wasm/tests/client_broker.rs @@ -0,0 +1,101 @@ +#![cfg(all(target_arch = "wasm32", feature = "broker"))] + +use mqtt5_wasm::{ + WasmBroker, WasmBrokerConfig, WasmConnectOptions, WasmMqttClient, WasmPublishOptions, + WasmSubscribeOptions, +}; +use std::cell::RefCell; +use std::rc::Rc; +use wasm_bindgen::prelude::*; +use wasm_bindgen::JsCast; +use wasm_bindgen_futures::JsFuture; +use wasm_bindgen_test::wasm_bindgen_test; + +async fn sleep(ms: i32) { + let promise = js_sys::Promise::new(&mut |resolve, _| { + let set_timeout = js_sys::Reflect::get(&js_sys::global(), &JsValue::from_str("setTimeout")) + .unwrap() + .unchecked_into::(); + set_timeout + .call2(&JsValue::NULL, &resolve, &JsValue::from(ms)) + .unwrap(); + }); + JsFuture::from(promise).await.unwrap(); +} + +fn publish_options(qos: u8) -> WasmPublishOptions { + let mut options = WasmPublishOptions::new(); + options.set_qos(qos); + options +} + +#[wasm_bindgen_test] +async fn wasm_client_against_wasm_broker() { + let mut config = WasmBrokerConfig::new(); + config.set_allow_anonymous(true); + let broker = WasmBroker::with_config(config).unwrap(); + + let subscriber = WasmMqttClient::new("sub".to_string()); + subscriber + .connect_message_port_with_options( + broker.create_client_port().unwrap(), + &WasmConnectOptions::new(), + ) + .await + .unwrap(); + let received = Rc::new(RefCell::new(Vec::<(String, u32)>::new())); + let sink = Rc::clone(&received); + let callback = Closure::::new( + move |topic: JsValue, payload: JsValue, _: JsValue| { + let length = js_sys::Uint8Array::new(&payload).length(); + sink.borrow_mut() + .push((topic.as_string().unwrap_or_default(), length)); + }, + ); + let mut subscribe_options = WasmSubscribeOptions::new(); + subscribe_options.set_qos(2); + subscriber + .subscribe_with_options( + "t/#", + callback.into_js_value().unchecked_into(), + &subscribe_options, + ) + .await + .unwrap(); + + let publisher = WasmMqttClient::new("pub".to_string()); + publisher + .connect_message_port_with_options( + broker.create_client_port().unwrap(), + &WasmConnectOptions::new(), + ) + .await + .unwrap(); + publisher.publish("t/0", b"").await.unwrap(); + publisher + .publish_with_options("t/1", b"a", &publish_options(1)) + .await + .unwrap(); + publisher + .publish_with_options("t/2", &vec![7u8; 200_000], &publish_options(2)) + .await + .unwrap(); + for i in 0..50u8 { + publisher + .publish_with_options(&format!("t/m{i}"), b"x", &publish_options(i % 3)) + .await + .unwrap(); + } + for _ in 0..200 { + if received.borrow().len() >= 53 { + break; + } + sleep(10).await; + } + let got = received.borrow().clone(); + assert_eq!(got.len(), 53, "{got:?}"); + assert!(got.contains(&("t/2".to_string(), 200_000))); + assert!(publisher.is_connected() && subscriber.is_connected()); + publisher.disconnect().await.unwrap(); + subscriber.disconnect().await.unwrap(); +} diff --git a/crates/mqtt5-wasm/tests/conformance_client.rs b/crates/mqtt5-wasm/tests/conformance_client.rs index 127dc30d..111a2a7e 100644 --- a/crates/mqtt5-wasm/tests/conformance_client.rs +++ b/crates/mqtt5-wasm/tests/conformance_client.rs @@ -14,7 +14,10 @@ use mqtt5_protocol::packet::{FixedHeader, MqttPacket, Packet}; use mqtt5_protocol::protocol::v5::properties::{PropertyId, PropertyValue}; use mqtt5_protocol::protocol::v5::reason_codes::ReasonCode; use mqtt5_protocol::QoS; -use mqtt5_wasm::{WasmConnectOptions, WasmMqttClient, WasmPublishOptions, WasmSubscribeOptions}; +use mqtt5_wasm::{ + WasmConnectOptions, WasmMqttClient, WasmPublishOptions, WasmReconnectOptions, + WasmSubscribeOptions, +}; use std::cell::{Cell, RefCell}; use std::rc::Rc; use wasm_bindgen::prelude::*; @@ -46,7 +49,7 @@ struct Frame { packet: Packet, } -fn take_frame(inbox: &mut Vec) -> Option { +fn take_frame(inbox: &mut Vec, protocol_version: &Cell) -> Option { let mut cursor = &inbox[..]; let header = FixedHeader::decode(&mut cursor).ok()?; let header_len = inbox.len() - cursor.len(); @@ -56,7 +59,16 @@ fn take_frame(inbox: &mut Vec) -> Option { } let first_byte = inbox[0]; let mut body = &inbox[header_len..total]; - let packet = Packet::decode_from_body(header.packet_type, &header, &mut body).unwrap(); + let packet = Packet::decode_from_body_with_version( + header.packet_type, + &header, + &mut body, + protocol_version.get(), + ) + .unwrap(); + if let Packet::Connect(connect) = &packet { + protocol_version.set(connect.protocol_version); + } inbox.drain(..total); Some(Frame { first_byte, packet }) } @@ -64,6 +76,7 @@ fn take_frame(inbox: &mut Vec) -> Option { struct FakeBroker { port: MessagePort, inbox: Rc>>, + protocol_version: Cell, closed: Rc>, on_message: Closure, on_close: Closure, @@ -91,6 +104,7 @@ impl FakeBroker { Self { port, inbox, + protocol_version: Cell::new(5), closed, on_message, on_close, @@ -109,7 +123,7 @@ impl FakeBroker { } fn take_frame(&self) -> Option { - take_frame(&mut self.inbox.borrow_mut()) + take_frame(&mut self.inbox.borrow_mut(), &self.protocol_version) } async fn next_frame(&self) -> Frame { @@ -256,6 +270,17 @@ fn value_recorder() -> (js_sys::Function, Rc>>) { (callback.into_js_value().unchecked_into(), values) } +type PairLog = Rc>>; + +fn pair_recorder() -> (js_sys::Function, PairLog) { + let values = Rc::new(RefCell::new(Vec::new())); + let sink = Rc::clone(&values); + let callback = Closure::::new(move |first, second| { + sink.borrow_mut().push((first, second)); + }); + (callback.into_js_value().unchecked_into(), values) +} + fn success() -> ConnAckPacket { ConnAckPacket::new(false, ReasonCode::Success) } @@ -309,7 +334,7 @@ fn publish_with( topic: &str, payload: &[u8], options: WasmPublishOptions, -) -> Outcome<()> { +) -> Outcome { let client = Rc::clone(client); let topic = topic.to_string(); let payload = payload.to_vec(); @@ -621,21 +646,60 @@ async fn mqtt_3_2_2_14_capabilities_refreshed_on_reconnect() { } #[wasm_bindgen_test] -async fn mqtt_3_2_2_11_maximum_qos_honoured() { +async fn mqtt_3_2_2_11_maximum_qos_zero_downgrades_and_reports() { let (client, broker) = connect_with(WasmConnectOptions::new(), success().with_maximum_qos(0)).await; let outcome = publish_with(&client, "a", b"x", publish_options(1)); - let qos1 = client.publish_qos1("a", b"x", noop()).await; - let qos2 = client.publish_qos2("a", b"x", noop()).await; - sleep(60).await; - while let Some(frame) = broker.take_frame() { - if let Packet::Publish(publish) = frame.packet { - assert_eq!(publish.qos, QoS::AtMostOnce, "PUBLISH exceeds Maximum QoS"); - } + let (qos1_callback, qos1_values) = pair_recorder(); + let (qos2_callback, qos2_values) = pair_recorder(); + let qos1 = client.publish_qos1("b", b"x", qos1_callback).await; + let qos2 = client.publish_qos2("c", b"x", qos2_callback).await; + let mut topics = Vec::new(); + for _ in 0..3 { + let publish = expect_publish(&broker).await; + assert_eq!(publish.qos, QoS::AtMostOnce, "PUBLISH exceeds Maximum QoS"); + topics.push(publish.topic_name); + } + topics.sort(); + assert_eq!(topics, ["a", "b", "c"]); + assert_eq!( + settle(&outcome).await.unwrap(), + 0, + "QoS used must be reported" + ); + assert_eq!(qos1.unwrap(), 0, "a QoS 0 publish has no packet identifier"); + assert_eq!(qos2.unwrap(), 0, "a QoS 0 publish has no packet identifier"); + for values in [&qos1_values, &qos2_values] { + let values = values.borrow(); + assert_eq!(values.len(), 1); + assert_eq!(values[0].0.as_f64(), Some(0.0)); + assert_eq!(values[0].1.as_f64(), Some(0.0), "QoS used must be reported"); } - assert!(settle(&outcome).await.is_err()); - assert!(qos1.is_err()); - assert!(qos2.is_err()); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_11_maximum_qos_one_downgrades_and_reports() { + let (client, broker) = + connect_with(WasmConnectOptions::new(), success().with_maximum_qos(1)).await; + let outcome = publish_with(&client, "a", b"x", publish_options(2)); + let first = expect_publish(&broker).await; + assert_eq!(first.qos, QoS::AtLeastOnce, "PUBLISH exceeds Maximum QoS"); + let (callback, values) = pair_recorder(); + let packet_id = client.publish_qos2("b", b"x", callback).await.unwrap(); + let second = expect_publish(&broker).await; + assert_eq!(second.qos, QoS::AtLeastOnce, "PUBLISH exceeds Maximum QoS"); + assert_eq!(second.packet_id, Some(packet_id)); + broker.send(&PubAckPacket::new(first.packet_id.unwrap())); + broker.send(&PubAckPacket::new(packet_id)); + assert_eq!( + settle(&outcome).await.unwrap(), + 1, + "QoS used must be reported" + ); + wait_for_count(&values, 1).await; + let values = values.borrow(); + assert_eq!(values[0].0.as_f64(), Some(0.0)); + assert_eq!(values[0].1.as_f64(), Some(1.0), "QoS used must be reported"); } #[wasm_bindgen_test] @@ -976,29 +1040,245 @@ async fn mqtt_3_2_2_4_resume_existing_session_still_rejects_clean_start() { ); } +const NOT_RESENT: &str = "indeterminate: may have been delivered; not resent because the new connection's limits do not allow it"; +const SESSION_LOST: &str = "indeterminate: session lost; may have been delivered"; +const CLEAN_START_DISCARD: &str = + "indeterminate: session discarded by clean start; may have been delivered"; +const SESSION_ENDED: &str = + "indeterminate: session ended with the connection; may have been delivered"; + +fn rejection_text(result: Result) -> String { + result + .expect_err("publish must be settled with an error") + .as_string() + .unwrap_or_default() +} + +async fn lose_connection(client: &WasmMqttClient, broker: &FakeBroker) { + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + sleep(50).await; + assert!(!client.is_connected()); +} + +async fn cycle_packet_ids(client: &WasmMqttClient, broker: &FakeBroker) -> Vec { + let mut issued = Vec::with_capacity(66_000); + while issued.len() < 65_600 { + let mut ids = Vec::with_capacity(2000); + for _ in 0..2000 { + ids.push(client.publish_qos1("t", b"", noop()).await.unwrap()); + } + let mut acks = Vec::with_capacity(ids.len() * 4); + for _ in &ids { + let publish = expect_publish(broker).await; + acks.extend(encode(&PubAckPacket::new(publish.packet_id.unwrap()))); + } + broker.send_raw(&acks); + sleep(20).await; + issued.extend(ids); + } + issued +} + #[wasm_bindgen_test] -async fn mqtt_3_2_2_5_session_state_discarded_on_session_present_zero() { - let client = Rc::new(WasmMqttClient::new("discard".to_string())); - let mut options = WasmConnectOptions::new(); - options.set_clean_start(false); - options.set_session_expiry_interval(Some(3600)); - let (broker, result, _) = open_session(&client, options, success()).await; +async fn mqtt_4_4_0_1_replay_skips_messages_the_new_limits_forbid_and_quarantines_qos2() { + let client = Rc::new(WasmMqttClient::new("replay-limits".to_string())); + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; result.unwrap(); + let mut retained = publish_options(1); + retained.set_retain(true); + let mut sent = Vec::new(); + let retained = publish_with(&client, "ret", b"r", retained); + sent.push(expect_publish(&broker).await); + let oversized = publish_with(&client, "big", &[0u8; 200], publish_options(1)); + sent.push(expect_publish(&broker).await); + let (qos2_callback, qos2_values) = pair_recorder(); + let quarantined = client + .publish_qos2("q2", b"q", qos2_callback) + .await + .unwrap(); + sent.push(expect_publish(&broker).await); + let conforming = publish_with(&client, "ok", b"1", publish_options(1)); + sent.push(expect_publish(&broker).await); + let released = publish_with(&client, "rel", b"2", publish_options(2)); + sent.push(expect_publish(&broker).await); + assert_eq!(sent[2].packet_id, Some(quarantined)); + let released_id = sent[4].packet_id.unwrap(); + broker.send(&PubRecPacket::new(released_id)); + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::PubRel(_) + )); + lose_connection(&client, &broker).await; + + let (broker, result, _) = open_session( + &client, + persistent_session_options(), + ConnAckPacket::new(true, ReasonCode::Success) + .with_retain_available(false) + .with_maximum_packet_size(64) + .with_maximum_qos(1), + ) + .await; + result.unwrap(); + let resent = broker.next_non_ping().await; + match &resent.packet { + Packet::Publish(publish) => { + assert_eq!(publish.topic_name, "ok"); + assert_eq!(publish.packet_id, sent[3].packet_id); + assert!(publish.dup); + } + other => panic!("expected the conforming PUBLISH, got {other:?}"), + } + match broker.next_non_ping().await.packet { + Packet::PubRel(pubrel) => assert_eq!(pubrel.packet_id, released_id), + other => panic!("expected the PUBREL to be resent, got {other:?}"), + } + broker.expect_silence(60).await; + + for outcome in [&retained, &oversized] { + let text = rejection_text(settle(outcome).await); + assert!(text.contains(NOT_RESENT), "rejection: {text}"); + } + wait_for_count(&qos2_values, 1).await; + let text = qos2_values.borrow()[0].0.as_string().unwrap_or_default(); + assert!(text.contains(NOT_RESENT), "callback value: {text}"); + + broker.send(&PubAckPacket::new(sent[3].packet_id.unwrap())); + broker.send(&PubCompPacket::new(released_id)); + assert_eq!(settle(&conforming).await.unwrap(), 1); + assert_eq!(settle(&released).await.unwrap(), 2); + + let issued = cycle_packet_ids(&client, &broker).await; + assert!( + !issued.contains(&quarantined), + "abandoned QoS 2 packet identifier {quarantined} reused while the session lasts" + ); + lose_connection(&client, &broker).await; + + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; + result.unwrap(); + let issued = cycle_packet_ids(&client, &broker).await; + assert!( + issued.contains(&quarantined), + "Session Present=0 must release quarantined packet identifier {quarantined}" + ); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_5_session_present_zero_resends_qos1_as_new_messages() { + let client = Rc::new(WasmMqttClient::new("requeue".to_string())); + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; + result.unwrap(); + let first = publish_with(&client, "a", b"1", publish_options(1)); + expect_publish(&broker).await; + let mut retained = publish_options(1); + retained.set_retain(true); + let retained = publish_with(&client, "r", b"2", retained); + expect_publish(&broker).await; + let (callback, values) = pair_recorder(); + client.publish_qos1("b", b"3", callback).await.unwrap(); + expect_publish(&broker).await; + lose_connection(&client, &broker).await; + + let (broker, result, _) = open_session( + &client, + persistent_session_options(), + success().with_retain_available(false), + ) + .await; + result.unwrap(); + let mut resent = Vec::new(); + for expected in ["a", "b"] { + let frame = broker.next_non_ping().await; + assert_eq!( + frame.first_byte & 0x08, + 0, + "a message sent on a new session has DUP=0" + ); + match frame.packet { + Packet::Publish(publish) => { + assert_eq!(publish.topic_name, expected); + assert!(!publish.dup); + resent.push(publish.packet_id.unwrap()); + } + other => panic!("expected PUBLISH {expected}, got {other:?}"), + } + } + broker.expect_silence(60).await; + let text = rejection_text(settle(&retained).await); + assert!(text.contains(SESSION_LOST), "rejection: {text}"); + assert!(text.contains("not resent"), "rejection: {text}"); + + for packet_id in &resent { + broker.send(&PubAckPacket::new(*packet_id)); + } + assert_eq!(settle(&first).await.unwrap(), 1); + wait_for_count(&values, 1).await; + let values = values.borrow(); + assert_eq!(values[0].0.as_f64(), Some(0.0)); + assert_eq!(values[0].1.as_f64(), Some(1.0)); +} + +#[wasm_bindgen_test] +async fn mqtt_3_1_2_4_clean_start_discards_kept_session_as_indeterminate() { + let client = Rc::new(WasmMqttClient::new("clean-discard".to_string())); + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; + result.unwrap(); + let qos1 = publish_with(&client, "a", b"1", publish_options(1)); + expect_publish(&broker).await; + let (callback, values) = pair_recorder(); + client.publish_qos2("b", b"2", callback).await.unwrap(); + expect_publish(&broker).await; + lose_connection(&client, &broker).await; + assert!(is_pending(&qos1)); + + let (broker, result, connect) = + open_session(&client, WasmConnectOptions::new(), success()).await; + result.unwrap(); + assert!(matches!(connect, Packet::Connect(connect) if connect.clean_start)); + broker.expect_silence(60).await; + let text = rejection_text(settle(&qos1).await); + assert!(text.contains(CLEAN_START_DISCARD), "rejection: {text}"); + wait_for_count(&values, 1).await; + let text = values.borrow()[0].0.as_string().unwrap_or_default(); + assert!(text.contains(CLEAN_START_DISCARD), "callback value: {text}"); +} + +#[wasm_bindgen_test] +async fn session_expiry_zero_connection_loss_settles_flights_as_indeterminate() { + let (client, broker) = connect_default().await; let pending = publish_with(&client, "a", b"1", publish_options(1)); + expect_publish(&broker).await; + lose_connection(&client, &broker).await; + let text = rejection_text(settle(&pending).await); + assert!(text.contains(SESSION_ENDED), "rejection: {text}"); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_5_session_present_zero_settles_qos2_as_indeterminate() { + let client = Rc::new(WasmMqttClient::new("qos2-lost".to_string())); + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; + result.unwrap(); + let unreleased = publish_with(&client, "q", b"1", publish_options(2)); + let (callback, values) = pair_recorder(); + client.publish_qos2("r", b"2", callback).await.unwrap(); + expect_publish(&broker).await; + let released_id = expect_publish(&broker).await.packet_id.unwrap(); + broker.send(&PubRecPacket::new(released_id)); assert!(matches!( broker.next_non_ping().await.packet, - Packet::Publish(_) + Packet::PubRel(_) )); - broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); - sleep(50).await; + lose_connection(&client, &broker).await; - let mut options = WasmConnectOptions::new(); - options.set_clean_start(false); - options.set_session_expiry_interval(Some(3600)); - let (broker, result, _) = open_session(&client, options, success()).await; + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; result.unwrap(); broker.expect_silence(60).await; - assert!(settle(&pending).await.is_err()); + let text = rejection_text(settle(&unreleased).await); + assert!(text.contains(SESSION_LOST), "rejection: {text}"); + wait_for_count(&values, 1).await; + let text = values.borrow()[0].0.as_string().unwrap_or_default(); + assert!(text.contains(SESSION_LOST), "callback value: {text}"); } #[wasm_bindgen_test] @@ -1415,6 +1695,8 @@ export class WsFake { } onConn(sock) { this.sock = sock; + this.closed = false; + this.buf = Buffer.alloc(0); let handshaken = false; sock.on('data', (chunk) => { this.buf = Buffer.concat([this.buf, chunk]); @@ -1432,7 +1714,7 @@ export class WsFake { } this.parse(); }); - sock.on('close', () => { this.closed = true; }); + sock.on('close', () => { if (this.sock === sock) this.closed = true; }); sock.on('error', () => {}); } parse() { @@ -1467,6 +1749,7 @@ export class WsFake { opcodeList() { return Uint8Array.from(this.opcodes); } protocolHeader() { return this.protocols; } isClosed() { return this.closed || !!this.closeReceived; } + dropClient() { if (this.sock) this.sock.destroy(); } stop() { try { if (this.sock) this.sock.destroy(); } catch (e) {} this.server.close(); } } "#)] @@ -1488,6 +1771,8 @@ extern "C" { fn protocol_header(this: &WsFake) -> String; #[wasm_bindgen(method, js_name = isClosed)] fn is_closed(this: &WsFake) -> bool; + #[wasm_bindgen(method, js_name = dropClient)] + fn drop_client(this: &WsFake); #[wasm_bindgen(method)] fn stop(this: &WsFake); } @@ -1495,13 +1780,14 @@ extern "C" { struct WsBroker { server: WsFake, inbox: RefCell>, + protocol_version: Cell, } impl WsBroker { async fn next_packet(&self) -> Packet { for _ in 0..400 { self.inbox.borrow_mut().extend(self.server.take_data()); - if let Some(frame) = take_frame(&mut self.inbox.borrow_mut()) { + if let Some(frame) = take_frame(&mut self.inbox.borrow_mut(), &self.protocol_version) { return frame.packet; } sleep(5).await; @@ -1530,6 +1816,7 @@ async fn connect_ws() -> (Rc, WsBroker) { let broker = WsBroker { server, inbox: RefCell::new(Vec::new()), + protocol_version: Cell::new(5), }; let client = Rc::new(WasmMqttClient::new("ws-client".to_string())); let url = format!("ws://127.0.0.1:{port}/mqtt"); @@ -1604,3 +1891,330 @@ async fn mqtt_6_0_0_2_websocket_coalesced_and_split_packets() { client.disconnect().await.unwrap(); broker.server.stop(); } + +fn persistent_session_options() -> WasmConnectOptions { + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options.set_session_expiry_interval(Some(3600)); + options +} + +fn resumable_without_expiry() -> WasmConnectOptions { + let mut options = WasmConnectOptions::new(); + options.set_clean_start(false); + options +} + +fn v311_options(clean_session: bool) -> WasmConnectOptions { + let mut options = WasmConnectOptions::new(); + options.set_protocol_version(4); + options.set_clean_start(clean_session); + options +} + +fn v311_connack(session_present: bool) -> ConnAckPacket { + ConnAckPacket::new_v311(session_present, ReasonCode::Success) +} + +async fn expect_publish(broker: &FakeBroker) -> PublishPacket { + match broker.next_non_ping().await.packet { + Packet::Publish(publish) => publish, + other => panic!("expected PUBLISH, got {other:?}"), + } +} + +async fn wait_until_connected(client: &WasmMqttClient) -> bool { + for _ in 0..400 { + if client.is_connected() { + return true; + } + sleep(5).await; + } + false +} + +#[wasm_bindgen_test] +async fn mqtt_3_1_2_4_v311_persistent_session_survives_connection_loss() { + let client = Rc::new(WasmMqttClient::new("v311-persistent".to_string())); + let (broker, result, _) = open_session(&client, v311_options(false), v311_connack(false)).await; + result.unwrap(); + let held = publish_with(&client, "a", b"1", publish_options(1)); + let first = expect_publish(&broker).await; + broker.send_raw(&[0xE0, 0x00]); + sleep(50).await; + assert!(!client.is_connected()); + assert!( + is_pending(&held), + "in-flight QoS 1 publish must stay in the session" + ); + + let (broker, result, connect) = + open_session(&client, v311_options(false), v311_connack(true)).await; + match connect { + Packet::Connect(connect) => { + assert_eq!(connect.protocol_version, 4); + assert!(!connect.clean_start); + } + other => panic!("expected CONNECT, got {other:?}"), + } + result.expect("CleanSession=0 reconnect to a persistent v3.1.1 session must succeed"); + let resent = expect_publish(&broker).await; + assert_eq!(resent.packet_id, first.packet_id); + assert!(resent.dup); + assert_eq!(resent.payload.as_ref(), b"1"); + broker.send(&PubAckPacket::new(resent.packet_id.unwrap())); + settle(&held).await.unwrap(); +} + +async fn ws_reconnect_clean_start(clean_session: bool) -> (bool, bool) { + let server = WsFake::new(); + let port = JsFuture::from(server.start()) + .await + .unwrap() + .as_f64() + .unwrap(); + let broker = WsBroker { + server, + inbox: RefCell::new(Vec::new()), + protocol_version: Cell::new(5), + }; + let client = Rc::new(WasmMqttClient::new("ws-v311".to_string())); + let mut reconnect = WasmReconnectOptions::new(); + reconnect.set_initial_delay_ms(10); + client.set_reconnect_options(&reconnect); + let url = format!("ws://127.0.0.1:{port}/mqtt"); + let connecting = { + let client = Rc::clone(&client); + spawn_outcome(async move { + client + .connect_with_options(&url, &v311_options(clean_session)) + .await + }) + }; + let first = match broker.next_packet().await { + Packet::Connect(connect) => connect.clean_start, + other => panic!("expected CONNECT, got {other:?}"), + }; + broker.server.send_binary(&encode(&v311_connack(false))); + settle(&connecting).await.unwrap(); + + broker.server.drop_client(); + let second = match broker.next_packet().await { + Packet::Connect(connect) => connect.clean_start, + other => panic!("expected reconnect CONNECT, got {other:?}"), + }; + broker.server.send_binary(&encode(&v311_connack(!second))); + assert!(wait_until_connected(&client).await, "reconnect failed"); + client.disconnect().await.unwrap(); + broker.server.stop(); + (first, second) +} + +#[wasm_bindgen_test] +async fn mqtt_3_1_2_4_v311_reconnect_keeps_clean_session_one() { + let (first, second) = ws_reconnect_clean_start(true).await; + assert!(first); + assert!( + second, + "reconnect must not turn a CleanSession=1 client into a persistent session" + ); +} + +#[wasm_bindgen_test] +async fn mqtt_3_1_2_4_v311_reconnect_keeps_clean_session_zero() { + let (first, second) = ws_reconnect_clean_start(false).await; + assert!(!first); + assert!(!second); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_3_2_server_session_expiry_extends_session() { + let client = Rc::new(WasmMqttClient::new("server-expiry-longer".to_string())); + let (broker, result, _) = open_session( + &client, + resumable_without_expiry(), + success().with_session_expiry_interval(300), + ) + .await; + result.unwrap(); + let held = publish_with(&client, "a", b"1", publish_options(1)); + let first = expect_publish(&broker).await; + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + sleep(50).await; + assert!( + is_pending(&held), + "session kept by the server's expiry must hold the flight" + ); + + let (broker, result, _) = open_session( + &client, + resumable_without_expiry(), + ConnAckPacket::new(true, ReasonCode::Success).with_session_expiry_interval(300), + ) + .await; + result.expect("the server's Session Expiry Interval keeps the session"); + let resent = expect_publish(&broker).await; + assert_eq!(resent.packet_id, first.packet_id); + assert!(resent.dup); + broker.send(&PubAckPacket::new(resent.packet_id.unwrap())); + settle(&held).await.unwrap(); +} + +#[wasm_bindgen_test] +async fn mqtt_3_2_2_3_2_server_session_expiry_zero_ends_session() { + let client = Rc::new(WasmMqttClient::new("server-expiry-zero".to_string())); + let (broker, result, _) = open_session( + &client, + persistent_session_options(), + success().with_session_expiry_interval(0), + ) + .await; + result.unwrap(); + let held = publish_with(&client, "a", b"1", publish_options(1)); + expect_publish(&broker).await; + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + let outcome = settle(&held).await; + assert!( + outcome.is_err(), + "a session the server ends with the connection must fail its flights" + ); + + let (broker, result, _) = open_session( + &client, + persistent_session_options(), + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + assert!(result.is_err(), "client holds no session after expiry 0"); + assert!(broker.wait_closed().await); +} + +#[wasm_bindgen_test] +async fn disconnect_settles_pending_acknowledgements_and_keeps_session() { + let client = Rc::new(WasmMqttClient::new("disconnect-settles".to_string())); + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; + result.unwrap(); + let pending = publish_with(&client, "a", b"1", publish_options(1)); + let (qos1_callback, qos1_values) = value_recorder(); + let (qos2_callback, qos2_values) = value_recorder(); + client.publish_qos1("b", b"2", qos1_callback).await.unwrap(); + client.publish_qos2("c", b"3", qos2_callback).await.unwrap(); + let mut sent = Vec::new(); + for _ in 0..3 { + sent.push(expect_publish(&broker).await); + } + + client.disconnect().await.unwrap(); + let outcome = settle(&pending).await; + let message = outcome + .expect_err("pending publish must be rejected by disconnect()") + .as_string() + .unwrap_or_default(); + assert!(message.contains("disconnected"), "rejection: {message}"); + wait_for_count(&qos1_values, 1).await; + wait_for_count(&qos2_values, 1).await; + for values in [&qos1_values, &qos2_values] { + let values = values.borrow(); + assert_eq!(values.len(), 1); + let text = values[0].as_string().unwrap_or_default(); + assert!(text.contains("disconnected"), "callback value: {text}"); + } + + let (broker, result, _) = open_session( + &client, + persistent_session_options(), + ConnAckPacket::new(true, ReasonCode::Success), + ) + .await; + result.expect("the kept session must resume on the same client instance"); + for original in &sent { + let resent = expect_publish(&broker).await; + assert_eq!(resent.packet_id, original.packet_id); + assert_eq!(resent.topic_name, original.topic_name); + assert!(resent.dup); + } + broker.send(&PubAckPacket::new(sent[0].packet_id.unwrap())); + broker.send(&PubAckPacket::new(sent[1].packet_id.unwrap())); + broker.send(&PubRecPacket::new(sent[2].packet_id.unwrap())); + match broker.next_non_ping().await.packet { + Packet::PubRel(pubrel) => assert_eq!(Some(pubrel.packet_id), sent[2].packet_id), + other => panic!("expected PUBREL, got {other:?}"), + } + broker.send(&PubCompPacket::new(sent[2].packet_id.unwrap())); + broker.expect_silence(50).await; + assert_eq!(qos1_values.borrow().len(), 1); + assert_eq!(qos2_values.borrow().len(), 1); +} + +#[wasm_bindgen_test] +async fn mqtt_3_3_4_7_resent_pubrels_above_lowered_receive_maximum_incur_quota_debt() { + let client = Rc::new(WasmMqttClient::new("quota-debt".to_string())); + let (broker, result, _) = open_session(&client, persistent_session_options(), success()).await; + result.unwrap(); + let first = publish_with(&client, "q2a", b"1", publish_options(2)); + let second = publish_with(&client, "q2b", b"2", publish_options(2)); + let third = publish_with(&client, "q1c", b"3", publish_options(1)); + let mut ids = Vec::new(); + for _ in 0..3 { + ids.push(expect_publish(&broker).await.packet_id.unwrap()); + } + broker.send(&PubRecPacket::new(ids[0])); + broker.send(&PubRecPacket::new(ids[1])); + for _ in 0..2 { + assert!(matches!( + broker.next_non_ping().await.packet, + Packet::PubRel(_) + )); + } + broker.send(&DisconnectPacket::new(ReasonCode::ServerShuttingDown)); + sleep(50).await; + + let (broker, result, _) = open_session( + &client, + persistent_session_options(), + ConnAckPacket::new(true, ReasonCode::Success).with_receive_maximum(1), + ) + .await; + result.unwrap(); + for expected in &ids[..2] { + match broker.next_non_ping().await.packet { + Packet::PubRel(pubrel) => assert_eq!(pubrel.packet_id, *expected), + other => panic!("expected PUBREL, got {other:?}"), + } + } + broker.expect_silence(50).await; + broker.send(&PubCompPacket::new(ids[0])); + broker.expect_silence(50).await; + broker.send(&PubCompPacket::new(ids[1])); + let resent = expect_publish(&broker).await; + assert_eq!(resent.packet_id, Some(ids[2])); + assert!(resent.dup); + broker.send(&PubAckPacket::new(ids[2])); + for outcome in [first, second, third] { + settle(&outcome).await.unwrap(); + } +} + +#[wasm_bindgen_test] +async fn v311_publish_internal_encoded_as_v311() { + let client = Rc::new(WasmMqttClient::new("v311-publish".to_string())); + let (broker, result, _) = open_session(&client, v311_options(true), v311_connack(false)).await; + result.unwrap(); + let publishing = { + let client = Rc::clone(&client); + spawn_outcome(async move { + client + .publish_internal("v311/topic", b"payload", QoS::AtLeastOnce) + .await + }) + }; + let frame = broker.next_non_ping().await; + let publish = match frame.packet { + Packet::Publish(publish) => publish, + other => panic!("expected PUBLISH, got {other:?}"), + }; + assert_eq!(publish.topic_name, "v311/topic"); + assert_eq!(publish.payload.as_ref(), b"payload"); + broker.send(&PubAckPacket::new(publish.packet_id.unwrap())); + settle(&publishing).await.unwrap(); +} diff --git a/crates/mqtt5/README.md b/crates/mqtt5/README.md index 15d046af..6115129e 100644 --- a/crates/mqtt5/README.md +++ b/crates/mqtt5/README.md @@ -415,17 +415,45 @@ AWS IoT features: - Client certificate loading from bytes (PEM/DER formats) - SDK compatibility: Subscribe method returns `(packet_id, qos)` tuple +### Publish Outcomes and the Offline Queue + +A publish on a live connection returns `PublishResult::Sent(Delivery)` once it is written (`QoS` 0) or acknowledged (`QoS` 1/2). `Delivery::qos_used()` reports the `QoS` actually used, which is lower than requested when the server's Maximum `QoS` forced a downgrade. + +A `QoS` 1/2 publish made while disconnected (with `queue_on_disconnect` enabled) returns `PublishResult::Queued(PublishHandle)`. The same happens when a `QoS` 1/2 publish was sent but the connection ended, the client disconnected, or no acknowledgement arrived within the acknowledgement wait: the message stays in flight with the session and is re-sent on the next connection. Either way it has not settled yet. The handle stays pending until the client reconnects (then it settles) or is dropped (then it resolves `Indeterminate(Abandoned)`); awaiting it yields exactly one `PublishOutcome`: + +- `Delivered(Delivery)`: acknowledged, or written unconfirmed when downgraded to `QoS` 0 +- `Rejected(PublishRejection)`: definitely not delivered (for example Retain Available 0 or Maximum Packet Size on the new connection) +- `Indeterminate(IndeterminateReason)`: may have been delivered (for example the session was lost before a `QoS` 2 PUBLISH was acknowledged with PUBREC, or an already-sent message no longer conforms to the resumed connection) + +```rust +use mqtt5::{MqttClient, PublishOutcome, PublishResult}; + +async fn send(client: &MqttClient) -> Result<(), Box> { + match client.publish_qos1("telemetry", b"data").await? { + PublishResult::Sent(delivery) => println!("sent at {:?}", delivery.qos_used()), + PublishResult::Queued(handle) => match handle.await { + PublishOutcome::Delivered(delivery) => println!("flushed at {:?}", delivery.qos_used()), + PublishOutcome::Rejected(reason) => println!("not delivered: {reason:?}"), + PublishOutcome::Indeterminate(reason) => println!("may have been delivered: {reason:?}"), + }, + } + Ok(()) +} +``` + +An offline retained publish is rejected immediately when the last connection reported Retain Available 0, and an offline publish that already exceeds the last known Maximum Packet Size is rejected immediately with `PacketTooLarge`. + ### Testing with Mock Client ```rust -use mqtt5::{MockMqttClient, MqttClientTrait, PublishResult, QoS}; +use mqtt5::{Delivery, MockMqttClient, MqttClientTrait, PublishResult, QoS}; #[tokio::test] async fn test_my_iot_function() { let mock = MockMqttClient::new("test-device"); mock.set_connect_response(Ok(())).await; - mock.set_publish_response(Ok(PublishResult::QoS1Or2 { packet_id: 123 })).await; + mock.set_publish_response(Ok(PublishResult::Sent(Delivery::AtLeastOnce { packet_id: 123 }))).await; my_iot_function(&mock).await.unwrap(); diff --git a/crates/mqtt5/src/broker/bridge/connection.rs b/crates/mqtt5/src/broker/bridge/connection.rs index b417b07d..710fbbc2 100644 --- a/crates/mqtt5/src/broker/bridge/connection.rs +++ b/crates/mqtt5/src/broker/bridge/connection.rs @@ -978,7 +978,7 @@ impl BridgeConnection { .publish_with_options(&remote_topic, payload, options) .await { - Ok(_) => { + Ok(crate::client::PublishResult::Sent(_)) => { debug!( bridge = %bridge_name_clone, topic = %remote_topic, @@ -987,6 +987,13 @@ impl BridgeConnection { messages_sent.fetch_add(1, Ordering::Relaxed); bytes_sent.fetch_add(payload_len as u64, Ordering::Relaxed); } + Ok(crate::client::PublishResult::Queued(_)) => { + debug!( + bridge = %bridge_name_clone, + topic = %remote_topic, + "publish not yet acknowledged; it stays in flight with the session" + ); + } Err(e) => { error!( bridge = %bridge_name_clone, diff --git a/crates/mqtt5/src/broker/client_handler/connect.rs b/crates/mqtt5/src/broker/client_handler/connect.rs index 2a8dfe1c..cb291387 100644 --- a/crates/mqtt5/src/broker/client_handler/connect.rs +++ b/crates/mqtt5/src/broker/client_handler/connect.rs @@ -1,4 +1,5 @@ use crate::broker::auth::EnhancedAuthStatus; +use crate::broker::router::SubscriptionRequest; use crate::broker::storage::{ClientSession, DynamicStorage, StorageBackend}; use crate::error::{MqttError, Result}; use crate::packet::auth::AuthPacket; @@ -513,16 +514,19 @@ impl ClientHandler { } self.router .subscribe( - connect.client_id.clone(), - topic_filter.clone(), - stored.qos, - stored.subscription_id, - stored.no_local, - stored.retain_as_published, - stored.retain_handling, - ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), - stored.change_only, - None, + SubscriptionRequest::new( + connect.client_id.clone(), + topic_filter.clone(), + stored.qos, + ) + .with_subscription_id(stored.subscription_id) + .with_no_local(stored.no_local) + .with_retain_as_published(stored.retain_as_published) + .with_retain_handling(stored.retain_handling) + .with_protocol_version( + ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), + ) + .with_change_only(stored.change_only), ) .await?; } diff --git a/crates/mqtt5/src/broker/client_handler/subscribe.rs b/crates/mqtt5/src/broker/client_handler/subscribe.rs index 24558e42..3519ec4c 100644 --- a/crates/mqtt5/src/broker/client_handler/subscribe.rs +++ b/crates/mqtt5/src/broker/client_handler/subscribe.rs @@ -15,7 +15,7 @@ use crate::validation::{parse_shared_subscription, topic_matches_filter, validat use crate::QoS; use tracing::{debug, warn}; -use crate::broker::router::{RoutableMessage, Subscribed, Unsubscribed}; +use crate::broker::router::{RoutableMessage, Subscribed, SubscriptionRequest, Unsubscribed}; use super::ClientHandler; @@ -54,16 +54,20 @@ impl ClientHandler { .router .subscribe_as( Some(self.generation), - client_id.clone(), - filter.filter.clone(), - QoS::from(granted_qos), - subscribe.properties.get_subscription_identifier(), - filter.options.no_local, - filter.options.retain_as_published, - filter.options.retain_handling as u8, - ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), - change_only, - flow_id, + SubscriptionRequest::new( + client_id.clone(), + filter.filter.clone(), + QoS::from(granted_qos), + ) + .with_subscription_id(subscribe.properties.get_subscription_identifier()) + .with_no_local(filter.options.no_local) + .with_retain_as_published(filter.options.retain_as_published) + .with_retain_handling(filter.options.retain_handling as u8) + .with_protocol_version( + ProtocolVersion::try_from(self.protocol_version).unwrap_or_default(), + ) + .with_change_only(change_only) + .with_flow_id(flow_id), ) .await?; let is_new = match outcome { diff --git a/crates/mqtt5/src/broker/router.rs b/crates/mqtt5/src/broker/router.rs index 857f66a2..3ab74b4f 100644 --- a/crates/mqtt5/src/broker/router.rs +++ b/crates/mqtt5/src/broker/router.rs @@ -22,6 +22,10 @@ use tracing::{debug, error, info, trace}; /// Upper bound on how long one publish may wait for slow subscribers' delivery channels. pub const ROUTE_BUDGET_MAX: Duration = Duration::from_secs(2); +fn default_route_deadline() -> Option { + cfg!(not(target_arch = "wasm32")).then(|| Instant::now() + ROUTE_BUDGET_MAX) +} + struct OutboundRateState { count: AtomicU32, window_start: parking_lot::Mutex, @@ -65,6 +69,83 @@ pub struct RoutableMessage { pub target_flow: Option, } +/// Parameters of one subscription registration passed to [`MessageRouter::subscribe`]. +#[derive(Debug, Clone)] +pub struct SubscriptionRequest { + pub client_id: String, + pub topic_filter: String, + pub qos: QoS, + pub subscription_id: Option, + pub no_local: bool, + pub retain_as_published: bool, + pub retain_handling: u8, + pub protocol_version: ProtocolVersion, + pub change_only: bool, + pub flow_id: Option, +} + +impl SubscriptionRequest { + /// Creates a request with no subscription identifier, all subscription options cleared, + /// `retain_handling` 0, MQTT v5.0, no change-only delivery and no flow. + #[must_use] + pub fn new(client_id: impl Into, topic_filter: impl Into, qos: QoS) -> Self { + Self { + client_id: client_id.into(), + topic_filter: topic_filter.into(), + qos, + subscription_id: None, + no_local: false, + retain_as_published: false, + retain_handling: 0, + protocol_version: ProtocolVersion::default(), + change_only: false, + flow_id: None, + } + } + + #[must_use] + pub fn with_subscription_id(mut self, subscription_id: Option) -> Self { + self.subscription_id = subscription_id; + self + } + + #[must_use] + pub fn with_no_local(mut self, no_local: bool) -> Self { + self.no_local = no_local; + self + } + + #[must_use] + pub fn with_retain_as_published(mut self, retain_as_published: bool) -> Self { + self.retain_as_published = retain_as_published; + self + } + + #[must_use] + pub fn with_retain_handling(mut self, retain_handling: u8) -> Self { + self.retain_handling = retain_handling; + self + } + + #[must_use] + pub fn with_protocol_version(mut self, protocol_version: ProtocolVersion) -> Self { + self.protocol_version = protocol_version; + self + } + + #[must_use] + pub fn with_change_only(mut self, change_only: bool) -> Self { + self.change_only = change_only; + self + } + + #[must_use] + pub fn with_flow_id(mut self, flow_id: Option) -> Self { + self.flow_id = flow_id; + self + } +} + /// Client subscription information #[derive(Debug, Clone)] pub struct Subscription { @@ -93,7 +174,7 @@ pub struct MessageRouter { #[cfg(not(target_arch = "wasm32"))] bridge_manager: Arc>>>, #[cfg(target_arch = "wasm32")] - wasm_bridge_callback: Arc>>, + wasm_bridge_callback: RwLock>, echo_suppression_key: Arc>>, outbound_rates: parking_lot::RwLock>, max_outbound_rate: AtomicU32, @@ -255,8 +336,7 @@ impl MessageRouter { #[cfg(not(target_arch = "wasm32"))] bridge_manager: Arc::new(RwLock::new(None)), #[cfg(target_arch = "wasm32")] - #[allow(clippy::arc_with_non_send_sync)] - wasm_bridge_callback: Arc::new(RwLock::new(None)), + wasm_bridge_callback: RwLock::new(None), echo_suppression_key: Arc::new(RwLock::new(None)), outbound_rates: parking_lot::RwLock::new(HashMap::new()), max_outbound_rate: AtomicU32::new(0), @@ -289,8 +369,7 @@ impl MessageRouter { #[cfg(not(target_arch = "wasm32"))] bridge_manager: Arc::new(RwLock::new(None)), #[cfg(target_arch = "wasm32")] - #[allow(clippy::arc_with_non_send_sync)] - wasm_bridge_callback: Arc::new(RwLock::new(None)), + wasm_bridge_callback: RwLock::new(None), echo_suppression_key: Arc::new(RwLock::new(None)), outbound_rates: parking_lot::RwLock::new(HashMap::new()), max_outbound_rate: AtomicU32::new(0), @@ -532,9 +611,6 @@ impl MessageRouter { } pub async fn cleanup_stale_subscriptions(&self) { - // Purge expired entries first (the storage backends do this for their own queues in - // cleanup_expired, but the router's fallback queues, used when persistence is off, have - // no backend sweep), then reclaim the now-empty ones. for queue in self.fallback_queues.handles() { queue.purge_expired(); } @@ -610,35 +686,8 @@ impl MessageRouter { /// /// # Errors /// Returns an error if subscription registration fails or `retain_handling` is invalid. - #[allow(clippy::too_many_arguments)] - pub async fn subscribe( - &self, - client_id: String, - topic_filter: String, - qos: QoS, - subscription_id: Option, - no_local: bool, - retain_as_published: bool, - retain_handling: u8, - protocol_version: ProtocolVersion, - change_only: bool, - flow_id: Option, - ) -> Result { - let outcome = self - .subscribe_as( - None, - client_id, - topic_filter, - qos, - subscription_id, - no_local, - retain_as_published, - retain_handling, - protocol_version, - change_only, - flow_id, - ) - .await?; + pub async fn subscribe(&self, request: SubscriptionRequest) -> Result { + let outcome = self.subscribe_as(None, request).await?; Ok(outcome == Subscribed::New) } @@ -647,21 +696,23 @@ impl MessageRouter { /// /// # Errors /// Returns an error when `retain_handling` is invalid. - #[allow(clippy::too_many_arguments)] pub async fn subscribe_as( &self, generation: Option, - client_id: String, - topic_filter: String, - qos: QoS, - subscription_id: Option, - no_local: bool, - retain_as_published: bool, - retain_handling: u8, - protocol_version: ProtocolVersion, - change_only: bool, - flow_id: Option, + request: SubscriptionRequest, ) -> Result { + let SubscriptionRequest { + client_id, + topic_filter, + qos, + subscription_id, + no_local, + retain_as_published, + retain_handling, + protocol_version, + change_only, + flow_id, + } = request; if retain_handling > 2 { return Err(crate::MqttError::ProtocolError(format!( "Invalid retain_handling value: {retain_handling} (must be 0, 1, or 2)" @@ -827,12 +878,8 @@ impl MessageRouter { /// Routes a publish message to all matching subscribers and forwards to bridges. pub async fn route_message(&self, publish: &PublishPacket, publishing_client_id: Option<&str>) { - self.route_message_with_deadline( - publish, - publishing_client_id, - Instant::now() + ROUTE_BUDGET_MAX, - ) - .await; + self.route_message_bounded(publish, publishing_client_id, default_route_deadline()) + .await; } /// Routes a publish message; `deadline` bounds the total time spent waiting on slow @@ -842,6 +889,16 @@ impl MessageRouter { publish: &PublishPacket, publishing_client_id: Option<&str>, deadline: Instant, + ) { + self.route_message_bounded(publish, publishing_client_id, Some(deadline)) + .await; + } + + async fn route_message_bounded( + &self, + publish: &PublishPacket, + publishing_client_id: Option<&str>, + deadline: Option, ) { #[cfg(feature = "opentelemetry")] { @@ -870,7 +927,7 @@ impl MessageRouter { publish: &PublishPacket, publishing_client_id: Option<&str>, ) { - let deadline = Instant::now() + ROUTE_BUDGET_MAX; + let deadline = default_route_deadline(); #[cfg(feature = "opentelemetry")] { use tracing::Instrument; @@ -895,7 +952,7 @@ impl MessageRouter { publish: &PublishPacket, publishing_client_id: Option<&str>, forward_to_bridges: bool, - deadline: Instant, + deadline: Option, ) { if forward_to_bridges { trace!("Routing message to topic: {}", publish.topic_name); @@ -1155,7 +1212,7 @@ impl MessageRouter { async fn execute_plan( plan: DeliveryPlan, publishing_client_id: Option<&str>, - deadline: Instant, + deadline: Option, ) { match plan { DeliveryPlan::Behind { @@ -1202,9 +1259,16 @@ impl MessageRouter { Self::queue_behind(&queue, routable.publish, &client_id, routable.target_flow); return; } - match tokio::time::timeout_at(deadline, lanes.qos1_tx.reserve()).await { - Ok(Ok(permit)) => permit.send(routable), - Ok(Err(_)) | Err(_) => Self::queue_behind( + let permit = match deadline { + Some(deadline) => tokio::time::timeout_at(deadline, lanes.qos1_tx.reserve()) + .await + .ok() + .and_then(std::result::Result::ok), + None => None, + }; + match permit { + Some(permit) => permit.send(routable), + None => Self::queue_behind( &queue, routable.publish, &client_id, @@ -1474,18 +1538,11 @@ mod tests { ) .await; router - .subscribe( - "client1".to_string(), - "test/+".to_string(), + .subscribe(SubscriptionRequest::new( + "client1", + "test/+", QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -1502,7 +1559,6 @@ mod tests { let mut rx1 = TestLanes::new(100); let mut rx2 = TestLanes::new(100); - // Register clients let (dtx1, _drx1) = tokio::sync::oneshot::channel(); let (dtx2, _drx2) = tokio::sync::oneshot::channel(); router @@ -1523,42 +1579,26 @@ mod tests { .await; router - .subscribe( - "client1".to_string(), - "test/+".to_string(), + .subscribe(SubscriptionRequest::new( + "client1", + "test/+", QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "client2".to_string(), - "test/data".to_string(), + .subscribe(SubscriptionRequest::new( + "client2", + "test/data", QoS::ExactlyOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); - // Publish message let publish = PublishPacket::new("test/data", &b"hello"[..], QoS::ExactlyOnce); router.route_message(&publish, None).await; - // Client 1 should receive with QoS 1 (downgraded) let rm1 = rx1.try_recv().unwrap(); assert_eq!(rm1.publish.topic_name, "test/data"); assert_eq!(rm1.publish.qos, QoS::AtLeastOnce); @@ -1582,18 +1622,7 @@ mod tests { ) .await; router - .subscribe( - "sub".to_string(), - "ids/#".to_string(), - QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + .subscribe(SubscriptionRequest::new("sub", "ids/#", QoS::AtLeastOnce)) .await .unwrap(); @@ -1620,18 +1649,11 @@ mod tests { ) .await; router - .subscribe( - "subscriber".to_string(), - "lock/#".to_string(), + .subscribe(SubscriptionRequest::new( + "subscriber", + "lock/#", QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -1655,18 +1677,11 @@ mod tests { .register_client(id.clone(), rx.lanes(), router.queue_handle(&id), dtx) .await; router - .subscribe( + .subscribe(SubscriptionRequest::new( id.clone(), - "lock/#".to_string(), + "lock/#", QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router.unregister_client(&id).await; @@ -1696,19 +1711,16 @@ mod tests { async fn test_retained_messages() { let router = MessageRouter::new(); - // Store retained message let mut publish = PublishPacket::new("test/status", &b"online"[..], QoS::AtMostOnce); publish.retain = true; router.route_message(&publish, None).await; assert_eq!(router.retained_count().await, 1); - // Get retained messages let retained = router.get_retained_messages("test/+").await; assert_eq!(retained.len(), 1); assert_eq!(retained[0].topic_name, "test/status"); - // Delete retained message let mut delete = PublishPacket::new("test/status", &b""[..], QoS::AtMostOnce); delete.retain = true; router.route_message(&delete, None).await; @@ -1723,7 +1735,6 @@ mod tests { let mut rx2 = TestLanes::new(100); let mut rx3 = TestLanes::new(100); - // Register three clients let (dtx1, _drx1) = tokio::sync::oneshot::channel(); let (dtx2, _drx2) = tokio::sync::oneshot::channel(); router @@ -1753,52 +1764,30 @@ mod tests { .await; router - .subscribe( - "client1".to_string(), - "$share/workers/test/data".to_string(), + .subscribe(SubscriptionRequest::new( + "client1", + "$share/workers/test/data", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "client2".to_string(), - "$share/workers/test/data".to_string(), + .subscribe(SubscriptionRequest::new( + "client2", + "$share/workers/test/data", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "client3".to_string(), - "$share/workers/test/data".to_string(), + .subscribe(SubscriptionRequest::new( + "client3", + "$share/workers/test/data", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); - // Publish 6 messages for i in 0..6 { let publish = PublishPacket::new( "test/data", @@ -1808,7 +1797,6 @@ mod tests { router.route_message(&publish, None).await; } - // Each client should receive exactly 2 messages let mut count1 = 0; let mut count2 = 0; let mut count3 = 0; @@ -1835,7 +1823,6 @@ mod tests { let mut rx2 = TestLanes::new(100); let mut rx3 = TestLanes::new(100); - // Register clients let (dtx1, _drx1) = tokio::sync::oneshot::channel(); router .register_client( @@ -1865,65 +1852,41 @@ mod tests { .await; router - .subscribe( - "shared1".to_string(), - "$share/group/test/+".to_string(), + .subscribe(SubscriptionRequest::new( + "shared1", + "$share/group/test/+", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "shared2".to_string(), - "$share/group/test/+".to_string(), + .subscribe(SubscriptionRequest::new( + "shared2", + "$share/group/test/+", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "regular".to_string(), - "test/+".to_string(), + .subscribe(SubscriptionRequest::new( + "regular", + "test/+", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); - // Publish message let publish = PublishPacket::new("test/data", &b"hello"[..], QoS::AtMostOnce); router.route_message(&publish, None).await; - // Regular subscriber should receive the message let regular_rm = rx3.try_recv().unwrap(); assert_eq!(®ular_rm.publish.payload[..], b"hello"); - // Only one of the shared subscribers should receive it let shared1_received = rx1.try_recv().is_ok(); let shared2_received = rx2.try_recv().is_ok(); - assert!(shared1_received ^ shared2_received); // XOR - exactly one should be true + assert!(shared1_received ^ shared2_received); } #[tokio::test] @@ -1952,33 +1915,19 @@ mod tests { .await; router - .subscribe( - "client1".to_string(), - "test/+".to_string(), + .subscribe(SubscriptionRequest::new( + "client1", + "test/+", QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "client2".to_string(), - "test/data".to_string(), + .subscribe(SubscriptionRequest::new( + "client2", + "test/data", QoS::ExactlyOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -2023,33 +1972,19 @@ mod tests { .await; router - .subscribe( - "client1".to_string(), - "test/echo".to_string(), + .subscribe(SubscriptionRequest::new( + "client1", + "test/echo", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "client2".to_string(), - "test/echo".to_string(), + .subscribe(SubscriptionRequest::new( + "client2", + "test/echo", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -2092,33 +2027,19 @@ mod tests { .await; router - .subscribe( - "client1".to_string(), - "test/echo".to_string(), + .subscribe(SubscriptionRequest::new( + "client1", + "test/echo", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "client2".to_string(), - "test/echo".to_string(), + .subscribe(SubscriptionRequest::new( + "client2", + "test/echo", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -2163,18 +2084,11 @@ mod tests { ) .await; router - .subscribe( - "sub1".to_string(), - "test/rate".to_string(), + .subscribe(SubscriptionRequest::new( + "sub1", + "test/rate", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -2209,18 +2123,11 @@ mod tests { ) .await; router - .subscribe( - "sub1".to_string(), - "test/rate".to_string(), + .subscribe(SubscriptionRequest::new( + "sub1", + "test/rate", QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -2265,18 +2172,11 @@ mod tests { ) .await; router - .subscribe( - "sub1".to_string(), - "test/rate".to_string(), + .subscribe(SubscriptionRequest::new( + "sub1", + "test/rate", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -2323,16 +2223,7 @@ mod tests { router .subscribe( - "c1".to_string(), - "sensor/#".to_string(), - QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - Some(42), + SubscriptionRequest::new("c1", "sensor/#", QoS::AtLeastOnce).with_flow_id(Some(42)), ) .await .unwrap(); @@ -2355,33 +2246,13 @@ mod tests { .await; router - .subscribe( - "c1".to_string(), - "data/#".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + .subscribe(SubscriptionRequest::new("c1", "data/#", QoS::AtMostOnce)) .await .unwrap(); router .subscribe( - "c1".to_string(), - "data/#".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - Some(7), + SubscriptionRequest::new("c1", "data/#", QoS::AtMostOnce).with_flow_id(Some(7)), ) .await .unwrap(); @@ -2406,49 +2277,20 @@ mod tests { .await; router - .subscribe( - "c1".to_string(), - "topic/a".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + .subscribe(SubscriptionRequest::new("c1", "topic/a", QoS::AtMostOnce)) .await .unwrap(); router .subscribe( - "c1".to_string(), - "topic/a".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - Some(10), + SubscriptionRequest::new("c1", "topic/a", QoS::AtMostOnce).with_flow_id(Some(10)), ) .await .unwrap(); router .subscribe( - "c1".to_string(), - "topic/b".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - Some(10), + SubscriptionRequest::new("c1", "topic/b", QoS::AtMostOnce).with_flow_id(Some(10)), ) .await .unwrap(); @@ -2477,32 +2319,14 @@ mod tests { router .subscribe( - "c1".to_string(), - "dup/test".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - Some(5), + SubscriptionRequest::new("c1", "dup/test", QoS::AtMostOnce).with_flow_id(Some(5)), ) .await .unwrap(); router .subscribe( - "c1".to_string(), - "dup/test".to_string(), - QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - Some(5), + SubscriptionRequest::new("c1", "dup/test", QoS::AtLeastOnce).with_flow_id(Some(5)), ) .await .unwrap(); diff --git a/crates/mqtt5/src/client/direct/ack.rs b/crates/mqtt5/src/client/direct/ack.rs index adce6b28..db84993f 100644 --- a/crates/mqtt5/src/client/direct/ack.rs +++ b/crates/mqtt5/src/client/direct/ack.rs @@ -51,6 +51,7 @@ pub(crate) struct AckRequest { #[derive(Default)] struct AckOrder { head: u64, + generation: u64, slots: VecDeque>, } @@ -62,6 +63,7 @@ impl AckOrder { fn discard(&mut self) { self.head += self.slots.len() as u64; + self.generation += 1; self.slots.clear(); } @@ -96,6 +98,19 @@ impl AckOrder { /// [`AckToken::reject`] consume it, making a double-acknowledgement a compile error. /// Dropping it without resolving emits a non-success acknowledgement and warns, so a /// forgotten token never wedges the window (`DeferredAckToken.tla`, obligation 7). +/// +/// # Head-of-line blocking +/// +/// PUBACK and PUBREC must leave the client in the order the PUBLISH packets arrived +/// (`[MQTT-4.6.0-2]`, `[MQTT-4.6.0-3]`). While deferred acknowledgement is enabled, an +/// unresolved token therefore holds back **every** later PUBACK/PUBREC on the connection, +/// not only those of other `subscribe_with_ack` deliveries: the automatic acknowledgements +/// of plain `subscribe` callbacks queue behind it too. Only `QoS` 1/2 messages received on a +/// QUIC data flow are acknowledged on their own stream and are not held back. Each held +/// acknowledgement keeps its message in the broker's in-flight window, so a token kept +/// for long can exhaust the client's Receive Maximum and stall all inbound `QoS` 1/2 +/// delivery. Resolve tokens promptly, in any order; the acknowledgements are released +/// in arrival order as soon as the earliest outstanding token is resolved. pub struct AckToken { seq: Option, packet_id: u16, @@ -210,9 +225,13 @@ impl AckDispatcher { let order = Arc::clone(&self.order); tokio::spawn(async move { while let Some(request) = rx.recv().await { - let ready = order.lock().release(request); + let (generation, ready) = { + let mut order = order.lock(); + let ready = order.release(request); + (order.generation, ready) + }; for next in ready { - Self::handle(next, &slot, &session).await; + Self::handle(next, generation, &order, &slot, &session).await; } } }); @@ -274,6 +293,8 @@ impl AckDispatcher { /// state is cleared, per `[MQTT-4.3.3-9]` (a later same-id PUBLISH is a new message). async fn handle( request: AckRequest, + generation: u64, + order: &Mutex, slot: &WriterSlot, session: &Arc>, ) { @@ -285,13 +306,12 @@ impl AckDispatcher { packet, release_inbound, } => { - Self::write(request.packet_id, packet, slot, session).await; + Self::write(request.packet_id, packet, generation, order, slot, session).await; if release_inbound { - session - .read() - .await - .acknowledge_inbound(request.packet_id) - .await; + let session = session.read().await; + if order.lock().generation == generation { + session.acknowledge_inbound(request.packet_id).await; + } } return; } @@ -313,6 +333,13 @@ impl AckDispatcher { let is_success = reason == ReasonCode::Success; { let session = session.read().await; + if order.lock().generation != generation { + debug!( + packet_id = request.packet_id, + "Dropping ack for a discarded session" + ); + return; + } match request.qos { QoS::AtMostOnce => {} QoS::ExactlyOnce if is_success => { @@ -328,19 +355,28 @@ impl AckDispatcher { } } - Self::write(request.packet_id, packet, slot, session).await; + Self::write(request.packet_id, packet, generation, order, slot, session).await; } async fn write( packet_id: u16, packet: Packet, + generation: u64, + order: &Mutex, slot: &WriterSlot, session: &Arc>, ) { if !ack_fits_server_maximum(session, &packet).await { return; } - let writer = slot.lock().await.clone(); + let writer = { + let current = slot.lock().await; + if order.lock().generation != generation { + debug!(packet_id, "Dropping ack for a discarded session"); + return; + } + current.clone() + }; let written = match &writer { Some(handle) => handle.lock().await.write_packet(packet).await.is_ok(), None => false, @@ -624,4 +660,72 @@ mod tests { "re-registering a filter replaces the earlier callback, as it does for exact filters" ); } + + #[tokio::test] + async fn ack_released_before_session_discard_is_dropped_after_it() { + use crate::client::direct::unified::UnifiedWriter; + use crate::session::state::AckResolution; + use tokio::io::AsyncReadExt; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let stream = tokio::net::TcpStream::connect(listener.local_addr().unwrap()) + .await + .unwrap(); + let (mut broker_side, _) = listener.accept().await.unwrap(); + let (_read_half, write_half) = stream.into_split(); + + let session = Arc::new(tokio::sync::RwLock::new(SessionState::new( + "t".to_string(), + SessionConfig::default(), + false, + ))); + let dispatcher = AckDispatcher::new(Arc::clone(&session)); + let seq = dispatcher.reserve(QoS::ExactlyOnce); + let (generation, released) = { + let mut order = dispatcher.order.lock(); + let released = order.release(AckRequest { + seq, + packet_id: 7, + qos: QoS::ExactlyOnce, + kind: AckKind::Ack, + }); + (order.generation, released) + }; + assert_eq!(released.len(), 1); + + dispatcher.discard_pending(); + session.read().await.clear_all_inbound_state().await; + *dispatcher.writer_slot.lock().await = Some(Arc::new(tokio::sync::Mutex::new( + UnifiedWriter::Tcp(write_half), + ))); + + for request in released { + AckDispatcher::handle( + request, + generation, + &dispatcher.order, + &dispatcher.writer_slot, + &session, + ) + .await; + } + + let fresh = session.read().await; + assert!( + !fresh.has_pubrec(7).await, + "a stale ack must not mark PUBREC sent on the fresh session" + ); + assert_eq!(fresh.get_resolution(7).await, AckResolution::Unresolved); + drop(fresh); + let mut buf = [0u8; 1]; + let written = tokio::time::timeout( + std::time::Duration::from_millis(200), + broker_side.read(&mut buf), + ) + .await; + assert!( + written.is_err(), + "a stale ack must not be written on the new connection" + ); + } } diff --git a/crates/mqtt5/src/client/direct/handlers.rs b/crates/mqtt5/src/client/direct/handlers.rs index c32dcba1..414d8792 100644 --- a/crates/mqtt5/src/client/direct/handlers.rs +++ b/crates/mqtt5/src/client/direct/handlers.rs @@ -399,10 +399,8 @@ async fn ack_qos2_inbound( #[cfg(feature = "transport-quic")] pub(super) async fn handle_incoming_packet_no_writer( packet: Packet, - callback_manager: &Arc, flow_id: Option, - keepalive_state: &Arc>, - codec_registry: Option<&Arc>, + handlers: &IncomingHandlers<'_>, ) -> Result<()> { match packet { Packet::Publish(mut publish) => { @@ -411,18 +409,19 @@ pub(super) async fn handle_incoming_packet_no_writer( "QoS > 0 publish received on unidirectional stream".to_string(), )); } - if let Some(registry) = codec_registry { + validate_inbound_publish(&mut publish, handlers.topic_aliases)?; + if let Some(registry) = handlers.codec_registry { let content_type = publish.properties.get_content_type(); let decoded = registry.decode_if_needed(&publish.payload, content_type.as_deref())?; publish.payload = decoded; } publish.stream_id = flow_id.map(|f| f.raw()); - let _ = callback_manager.dispatch(&publish); + let _ = handlers.callback_manager.dispatch(&publish); Ok(()) } Packet::PingResp => { - keepalive_state.lock().record_pong_received(); + handlers.keepalive_state.lock().record_pong_received(); Ok(()) } Packet::Disconnect(disconnect) => { diff --git a/crates/mqtt5/src/client/direct/keepalive.rs b/crates/mqtt5/src/client/direct/keepalive.rs index e239dd82..6808c7f5 100644 --- a/crates/mqtt5/src/client/direct/keepalive.rs +++ b/crates/mqtt5/src/client/direct/keepalive.rs @@ -11,6 +11,7 @@ use std::sync::Arc; use tokio::time::Duration; use super::unified::UnifiedWriter; +use crate::session::flow_control::FlowControlManager; #[cfg(feature = "transport-quic")] use crate::session::SessionState; @@ -78,6 +79,10 @@ pub(super) struct ConnectionLifecycle { pub(super) connection_epoch: u64, pub(super) current_connection_epoch: Arc, pub(super) callbacks: Arc>>, + pub(super) closing: Arc, + pub(super) alive: Arc>, + pub(super) flow: Arc>, + pub(super) reader_task: Arc>>, } impl ConnectionLifecycle { @@ -85,15 +90,34 @@ impl ConnectionLifecycle { owns_current_connection(self.connection_epoch, &self.current_connection_epoch) } + pub(super) fn begin_close(&self) { + self.closing.store(true, Ordering::SeqCst); + } + + #[cfg(feature = "transport-quic")] + pub(super) fn is_closing(&self) -> bool { + self.closing.load(Ordering::SeqCst) + || !self.connected.load(Ordering::SeqCst) + || !self.owns_current_connection() + } + pub(super) async fn end(&self, reason: DisconnectReason) { + self.alive.send_replace(false); if mark_disconnected_if_current( &self.connected, self.connection_epoch, &self.current_connection_epoch, ) { + self.flow.read().await.close_send_quota(); fire_connection_event(&self.callbacks, ConnectionEvent::Disconnected { reason }).await; } } + + fn stop_reader(&self) { + if let Some(reader) = self.reader_task.lock().take() { + reader.abort(); + } + } } pub(super) const PINGREQ_LOG_INTERVAL: u32 = 20; @@ -149,9 +173,14 @@ pub(super) async fn keepalive_task_with_writer( let timed_out = keepalive_state.lock().is_timeout(timeout_duration); if timed_out { tracing::error!("Keepalive timeout - no PINGRESP received"); - super::reader::close_connection(&writer, &crate::error::MqttError::KeepAliveTimeout) - .await; + super::reader::close_connection( + &writer, + &lifecycle, + &crate::error::MqttError::KeepAliveTimeout, + ) + .await; lifecycle.end(DisconnectReason::KeepAliveTimeout).await; + lifecycle.stop_reader(); break; } @@ -180,6 +209,7 @@ pub(super) async fn keepalive_task_with_writer( lifecycle .end(DisconnectReason::NetworkError(e.to_string())) .await; + lifecycle.stop_reader(); break; } Err(_) => { @@ -189,6 +219,7 @@ pub(super) async fn keepalive_task_with_writer( "PINGREQ send timed out".to_string(), )) .await; + lifecycle.stop_reader(); break; } } diff --git a/crates/mqtt5/src/client/direct/mod.rs b/crates/mqtt5/src/client/direct/mod.rs index 7687896e..e2f2f040 100644 --- a/crates/mqtt5/src/client/direct/mod.rs +++ b/crates/mqtt5/src/client/direct/mod.rs @@ -8,21 +8,25 @@ mod keepalive; mod outbound; mod reader; mod replay; +mod tracking; mod unified; pub use ack::AckToken; pub(crate) use ack::{AckCallbackManager, AckDispatcher}; use parking_lot::Mutex; -use std::collections::{HashMap, VecDeque}; +use std::collections::HashMap; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::Arc; -use tokio::sync::oneshot; +use tokio::sync::{oneshot, watch}; use tokio::task::JoinHandle; use tokio::time::Duration; use crate::callback::{CallbackId, CallbackManager}; use crate::client::auth_handler::{AuthHandler, AuthResponse}; +use crate::client::publish_outcome::{ + Delivery, IndeterminateReason, PublishHandle, PublishOutcome, PublishRejection, PublishResult, +}; use crate::error::{MqttError, Result}; use crate::packet::auth::AuthPacket; use crate::packet::connect::ConnectPacket; @@ -40,7 +44,7 @@ use crate::session::state::OutboundReplay; use crate::session::subscription::Subscription; use crate::session::SessionState; use crate::transport::{PacketIo, PacketWriter, TransportType}; -use crate::types::{ConnectOptions, ConnectResult, PublishOptions, PublishResult}; +use crate::types::{ConnectOptions, ConnectResult, PublishOptions}; use crate::QoS; #[cfg(feature = "opentelemetry")] @@ -64,7 +68,8 @@ use keepalive::{keepalive_task_with_writer, KeepaliveState}; #[cfg(feature = "transport-quic")] use reader::quic_stream_acceptor_task; use reader::{packet_reader_task_with_responses, PacketReaderContext}; -use replay::{PublishPolicy, SessionReplay}; +use replay::{ConnectionLink, OfflineQueue, PublishPolicy, QueuedPublish, SessionReplay}; +use tracking::{Completion, IdReservation, OutboundIds, OutcomeTracker, SharedIds, SharedOutcomes}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum AutomaticReconnectLifecycle { @@ -81,33 +86,62 @@ pub(crate) enum SubscriptionPersistence { pub(crate) type StoredSubscription = (String, SubscriptionOptions, Option, CallbackId); pub(crate) type StoredSubscriptions = Arc>>; pub(crate) type ConnectionEpoch = Arc; -pub(crate) type PendingAcks = Arc>>>; +const ACKNOWLEDGEMENT_WAIT: Duration = Duration::from_secs(10); + +#[derive(Debug)] +pub(crate) struct ReadyPublish { + request: PublishPacket, + packet: PublishPacket, + reservation: Option, + epoch: u64, +} + +impl ReadyPublish { + pub(crate) fn packet_id(&self) -> Option { + self.packet.packet_id + } +} #[derive(Debug)] pub(crate) enum StagedPublish { - Queued(PublishResult), - Ready(PublishPacket), + Queued(PublishHandle), + Ready(Box), } -pub(crate) struct PublishAck { - rx: oneshot::Receiver, - packet_id: u16, - pending: PendingAcks, +pub(crate) enum Transmitted { + Sent(Delivery), + InFlight(InFlight), + Detached(PublishHandle), + Restaged(Box), } -impl PublishAck { - pub(crate) async fn wait(self) -> Result<()> { - match tokio::time::timeout(Duration::from_secs(10), self.rx).await { - Ok(Ok(reason_code)) if reason_code.is_error() => { - Err(MqttError::PublishFailed(reason_code)) +pub(crate) struct InFlight { + handle: PublishHandle, + link: watch::Receiver, +} + +impl InFlight { + pub(crate) async fn settle(self) -> Result { + let Self { handle, mut link } = self; + let outcome = handle.clone().outcome(); + tokio::select! { + biased; + outcome = outcome => match outcome { + PublishOutcome::Delivered(delivery) => Ok(PublishResult::Sent(delivery)), + PublishOutcome::Rejected(PublishRejection::Refused(reason_code)) => { + Err(MqttError::PublishFailed(reason_code)) + } + PublishOutcome::Rejected(_) | PublishOutcome::Indeterminate(_) => { + Ok(PublishResult::Queued(handle)) + } + }, + _ = link.wait_for(|alive| !*alive) => { + tracing::debug!("Connection ended before the acknowledgement; publish stays in flight"); + Ok(PublishResult::Queued(handle)) } - Ok(Ok(_)) => Ok(()), - Ok(Err(_)) => Err(MqttError::ProtocolError( - "Acknowledgment channel closed".to_string(), - )), - Err(_) => { - self.pending.lock().remove(&self.packet_id); - Err(MqttError::Timeout) + () = tokio::time::sleep(ACKNOWLEDGEMENT_WAIT) => { + tracing::debug!("Acknowledgement not received in time; publish stays in flight"); + Ok(PublishResult::Queued(handle)) } } } @@ -145,13 +179,16 @@ pub struct DirectClientInner { pub packet_id_generator: PacketIdGenerator, pub pending_subacks: Arc>>>, pub pending_unsubacks: Arc>>>, - pub pending_pubacks: PendingAcks, - pub pending_pubcomps: PendingAcks, pub reconnect_attempt: u32, pub last_address: Option, pub automatic_reconnect_lifecycle: AutomaticReconnectLifecycle, pub server_redirect: Option, - pub queued_messages: Arc>>, + pub queued_messages: Arc>, + outbound_ids: SharedIds, + publish_outcomes: SharedOutcomes, + outbound_transfer: Arc>, + connection_alive: Option>>, + send_flow: Arc>, pub stored_subscriptions: StoredSubscriptions, pub stored_ack_subscriptions: StoredSubscriptions, pub queue_on_disconnect: bool, @@ -170,11 +207,13 @@ pub struct DirectClientInner { impl DirectClientInner { pub fn new(options: ConnectOptions) -> Self { - let session = Arc::new(tokio::sync::RwLock::new(SessionState::new( + let session_state = SessionState::new( options.client_id.clone(), options.session_config.clone(), options.clean_start, - ))); + ); + let send_flow = Arc::clone(session_state.flow_control()); + let session = Arc::new(tokio::sync::RwLock::new(session_state)); let queue_on_disconnect = !options.clean_start; let auth_method = options.properties.authentication_method.clone(); @@ -211,13 +250,16 @@ impl DirectClientInner { packet_id_generator: PacketIdGenerator::new(), pending_subacks: Arc::new(Mutex::new(HashMap::new())), pending_unsubacks: Arc::new(Mutex::new(HashMap::new())), - pending_pubacks: Arc::new(Mutex::new(HashMap::new())), - pending_pubcomps: Arc::new(Mutex::new(HashMap::new())), reconnect_attempt: 0, last_address: None, automatic_reconnect_lifecycle: AutomaticReconnectLifecycle::Armed, server_redirect: None, - queued_messages: Arc::new(Mutex::new(VecDeque::new())), + queued_messages: Arc::new(Mutex::new(OfflineQueue::default())), + outbound_ids: Arc::new(Mutex::new(OutboundIds::default())), + publish_outcomes: Arc::new(Mutex::new(OutcomeTracker::default())), + outbound_transfer: Arc::new(tokio::sync::Mutex::new(())), + connection_alive: None, + send_flow, stored_subscriptions: Arc::new(Mutex::new(Vec::new())), stored_ack_subscriptions: Arc::new(Mutex::new(Vec::new())), queue_on_disconnect, @@ -289,6 +331,10 @@ impl DirectClientInner { "resetting connection runtime" ); self.set_connected(false); + if let Some(alive) = self.connection_alive.take() { + alive.send_replace(false); + } + drop(self.outbound_transfer.lock().await); self.stop_background_tasks().await; self.keepalive_state.lock().reset(); @@ -488,14 +534,14 @@ impl DirectClientInner { let effective_flow_headers = split.flow_headers_enabled && split.negotiated_mqtt_next; self.quic_stream_manager = Some(Arc::new( - QuicStreamManager::new(conn_arc, split.strategy) + QuicStreamManager::new(Arc::clone(&conn_arc), split.strategy) .with_flow_headers(effective_flow_headers) .with_flow_expire_interval(split.flow_expire_interval) .with_flow_flags(split.flow_flags), )); ( UnifiedReader::quic(split.recv, protocol_version), - UnifiedWriter::Quic(split.send), + UnifiedWriter::QuicControl(split.send, conn_arc), ) } }; @@ -508,6 +554,7 @@ impl DirectClientInner { .await; let replay_writer = Arc::downgrade(&writer_arc); self.writer = Some(writer_arc); + self.connection_alive = Some(Arc::new(watch::channel(true).0)); self.set_connected(true); tracing::debug!("Starting background tasks (packet reader and keepalive)"); @@ -515,15 +562,10 @@ impl DirectClientInner { tracing::debug!("Background tasks started successfully"); if let Some(slots) = replay_slots { - let replay = SessionReplay { - items: replay_items, - slots, - session: Arc::clone(&self.session), - writer: replay_writer, - queued: Arc::clone(&self.queued_messages), - policy: self.publish_policy(), - }; - tokio::spawn(replay.run()); + tokio::spawn( + self.session_replay(replay_items, slots, replay_writer, connection_epoch) + .run(), + ); } Ok(ConnectResult { @@ -531,6 +573,35 @@ impl DirectClientInner { }) } + fn session_replay( + &self, + items: Vec, + slots: Arc, + writer: std::sync::Weak>, + epoch: u64, + ) -> SessionReplay { + SessionReplay { + items, + slots, + session: Arc::clone(&self.session), + writer, + queued: Arc::clone(&self.queued_messages), + policy: self.publish_policy(), + outcomes: Arc::clone(&self.publish_outcomes), + ids: Arc::clone(&self.outbound_ids), + link: ConnectionLink { + epoch, + current_epoch: Arc::clone(&self.connection_epoch), + connected: Arc::clone(&self.connected), + transfer: Arc::clone(&self.outbound_transfer), + alive: self + .connection_alive + .as_ref() + .map_or_else(|| watch::channel(false).1, |alive| alive.subscribe()), + }, + } + } + fn holds_session_state(&self) -> bool { !self.options.clean_start && (self.connection_epoch.load(Ordering::SeqCst) > 0 @@ -582,9 +653,11 @@ impl DirectClientInner { } async fn discard_session_state(&self) { + let session = self.session.write().await; self.ack_dispatcher.discard_pending(); - let session = self.session.read().await; + let unacknowledged = session.outbound_replay().await; session.discard_outbound_state().await; + self.requeue_lost_session(unacknowledged); session.flow_control().read().await.clear_inbound().await; if session.clear_all_inbound_state().await { if self.options.deferred_ack { @@ -601,6 +674,71 @@ impl DirectClientInner { } } + fn requeue_lost_session(&self, unacknowledged: Vec) { + self.outbound_ids.lock().release_quarantine(); + let resume_requested = !self.options.clean_start; + let lost = if resume_requested { + IndeterminateReason::SessionLost + } else { + IndeterminateReason::SessionDiscarded + }; + let mut outcomes = self.publish_outcomes.lock(); + let mut resend = Vec::new(); + for item in unacknowledged { + let (packet_id, requeue) = match item { + OutboundReplay::Publish(publish) => { + let requeue = publish.qos == QoS::AtLeastOnce && resume_requested; + (publish.packet_id, requeue.then_some(publish)) + } + OutboundReplay::PubRel(packet_id) => { + if let Some(completion) = outcomes.take(packet_id) { + tracing::debug!( + packet_id, + "Session not resumed after PUBREC; the server owns the message" + ); + completion.delivered(Delivery::ExactlyOnce { packet_id }); + } + continue; + } + }; + let Some(packet_id) = packet_id else { + continue; + }; + let completion = outcomes.take(packet_id); + let reservation = requeue + .as_ref() + .and_then(|_| IdReservation::claim(&self.outbound_ids, packet_id)); + if let (Some(publish), Some(reservation)) = (requeue, reservation) { + tracing::debug!( + packet_id, + "Session not resumed; unacknowledged QoS 1 PUBLISH re-queued" + ); + let completion = completion.map(|mut completion| { + completion.mark_resent(); + completion + }); + resend.push(QueuedPublish::new( + PublishPacket { + dup: false, + ..publish + }, + reservation, + completion, + )); + } else if let Some(completion) = completion { + tracing::warn!( + packet_id, + ?lost, + "Session not resumed; unacknowledged outbound exchange dropped" + ); + completion.indeterminate(lost); + } + } + outcomes.abandon_all(lost); + drop(outcomes); + self.queued_messages.lock().push_front_in_order(resend); + } + async fn adopt_assigned_client_identifier( &mut self, connack: &crate::packet::connack::ConnAckPacket, @@ -741,6 +879,14 @@ impl DirectClientInner { /// Returns an error if the operation fails pub async fn disconnect_with_packet(&mut self, send_disconnect: bool) -> Result<()> { if !self.is_connected() { + self.reset_connection_runtime(b"disconnect").await; + self.session + .read() + .await + .flow_control() + .read() + .await + .close_send_quota(); return Err(MqttError::NotConnected); } @@ -770,17 +916,21 @@ impl DirectClientInner { /// # Errors /// - /// Returns `PacketTooLarge` when the message already exceeds the last known - /// negotiated maximum packet size, so a `QoS` 1/2 publish is rejected at - /// enqueue time instead of being acknowledged with `Ok` and then silently - /// dropped when the queue is flushed on reconnect. A packet identifier is - /// allocated only after the size check passes. + /// Returns `RetainNotSupported` or `PacketTooLarge` when the message already + /// violates the last known Retain Available or negotiated maximum packet size, so + /// the publish is rejected at enqueue time instead of being accepted and then + /// rejected when the queue is flushed on reconnect. A packet identifier is + /// allocated only after these checks pass. async fn queue_publish_message( &self, topic: String, payload: Vec, options: &PublishOptions, - ) -> Result { + ) -> Result { + if options.retain && !self.server_retain_available.load(Ordering::SeqCst) { + return Err(MqttError::RetainNotSupported); + } + let mut publish = self .with_aliased_topic(PublishPacket { topic_name: topic, @@ -797,10 +947,15 @@ impl DirectClientInner { self.check_publish_size(&publish).await?; - let packet_id = self.allocate_packet_id().await?; - publish.packet_id = Some(packet_id); - self.queued_messages.lock().push_back(publish); - Ok(PublishResult::QoS1Or2 { packet_id }) + let reservation = self.allocate_packet_id().await?; + publish.packet_id = Some(reservation.packet_id()); + let (completion, handle) = Completion::new(); + self.queued_messages.lock().push_back(QueuedPublish::new( + publish, + reservation, + Some(completion), + )); + Ok(handle) } async fn with_aliased_topic(&self, mut publish: PublishPacket) -> Result { @@ -818,21 +973,24 @@ impl DirectClientInner { Ok(publish) } - async fn allocate_packet_id(&self) -> Result { - self.session - .read() - .await - .allocate_packet_id(&self.packet_id_generator, |packet_id| { - self.pending_subacks.lock().contains_key(&packet_id) - || self.pending_unsubacks.lock().contains_key(&packet_id) - || self - .queued_messages - .lock() - .iter() - .any(|queued| queued.packet_id == Some(packet_id)) - }) - .await - .ok_or(MqttError::PacketIdExhausted) + async fn allocate_packet_id(&self) -> Result { + for _ in 0..u16::MAX { + let packet_id = self + .session + .read() + .await + .allocate_packet_id(&self.packet_id_generator, |packet_id| { + self.pending_subacks.lock().contains_key(&packet_id) + || self.pending_unsubacks.lock().contains_key(&packet_id) + || self.outbound_ids.lock().holds(packet_id) + }) + .await + .ok_or(MqttError::PacketIdExhausted)?; + if let Some(reservation) = IdReservation::claim(&self.outbound_ids, packet_id) { + return Ok(reservation); + } + } + Err(MqttError::PacketIdExhausted) } /// # Errors @@ -856,22 +1014,6 @@ impl DirectClientInner { self.check_packet_fits(packet).await } - fn setup_publish_acknowledgment(&self, qos: QoS, packet_id: Option) -> Option { - let pending = match qos { - QoS::AtMostOnce => return None, - QoS::AtLeastOnce => &self.pending_pubacks, - QoS::ExactlyOnce => &self.pending_pubcomps, - }; - let packet_id = packet_id?; - let (tx, rx) = oneshot::channel(); - pending.lock().insert(packet_id, tx); - Some(PublishAck { - rx, - packet_id, - pending: Arc::clone(pending), - }) - } - pub(super) async fn release_outbound_quota( session: &Arc>, packet_id: Option, @@ -900,7 +1042,8 @@ impl DirectClientInner { payload: Vec, options: PublishOptions, ) -> Result { - outbound::check_publish(&topic, &options)?; + let protocol_version = self.options.protocol_version.as_u8(); + outbound::check_publish(&topic, &options, protocol_version)?; if !self.is_connected() && self.queue_on_disconnect && options.qos != QoS::AtMostOnce { return self @@ -920,69 +1063,122 @@ impl DirectClientInner { return Err(MqttError::NotConnected); } - if let Some(alias) = options.properties.topic_alias { - let session = self.session.read().await; - let aliases = session.topic_alias_out().read().await; - outbound::check_topic_alias(&aliases, &topic, alias)?; - } let (final_payload, properties) = self.encode_payload(payload, &options)?; - - let mut publish = self.publish_policy().conform(PublishPacket { + let request = PublishPacket { topic_name: topic, payload: final_payload, qos: options.qos, retain: options.retain, dup: false, - packet_id: (options.qos != QoS::AtMostOnce).then_some(SIZE_PROBE_PACKET_ID), - properties, - protocol_version: self.options.protocol_version.as_u8(), + packet_id: None, + properties: if protocol_version == 5 { + properties + } else { + Properties::default() + }, + protocol_version, stream_id: None, - })?; + }; + self.conform_to_connection(request) + .await + .map(StagedPublish::Ready) + } - self.check_publish_size(&publish).await?; + async fn conform_to_connection(&self, request: PublishPacket) -> Result> { + if !self.is_connected() { + return Err(MqttError::NotConnected); + } - if publish.qos != QoS::AtMostOnce { - publish.packet_id = Some(self.allocate_packet_id().await?); + if let Some(alias) = request.topic_alias() { + let session = self.session.read().await; + let aliases = session.topic_alias_out().read().await; + outbound::check_topic_alias(&aliases, &request.topic_name, alias)?; } - Ok(StagedPublish::Ready(publish)) + let mut packet = self.publish_policy().conform(PublishPacket { + packet_id: (request.qos != QoS::AtMostOnce).then_some(SIZE_PROBE_PACKET_ID), + ..request.clone() + })?; + + self.check_publish_size(&packet).await?; + + let reservation = if packet.qos == QoS::AtMostOnce { + None + } else { + let reservation = self.allocate_packet_id().await?; + packet.packet_id = Some(reservation.packet_id()); + Some(reservation) + }; + + Ok(Box::new(ReadyPublish { + request, + packet, + reservation, + epoch: self.connection_epoch.load(Ordering::SeqCst), + })) } pub(crate) async fn transmit_publish( &self, - publish: PublishPacket, - ) -> Result> { - let qos = publish.qos; - let packet_id = publish.packet_id; + ready: Box, + claim: Option, + ) -> Result { + let packet_id = ready.packet_id(); let flow = Arc::clone(self.session.read().await.flow_control()); + let quota_generation = flow.read().await.quota_generation(); + let claimed = packet_id.filter(|_| claim == Some(quota_generation)); if !self.is_connected() { - if let Some(pid) = packet_id { + if let Some(pid) = claimed { Self::release_send_quota(&flow, pid).await; } return Err(MqttError::NotConnected); } - if qos != QoS::AtMostOnce { - let stored = match self.with_aliased_topic(publish.clone()).await { - Ok(stored) => { - self.session - .read() - .await - .store_unacked_publish(stored) - .await - } - Err(e) => Err(e), - }; - if let Err(e) = stored { - if let Some(pid) = packet_id { - Self::release_send_quota(&flow, pid).await; - } - return Err(e); + if ready.epoch != self.connection_epoch.load(Ordering::SeqCst) { + if let Some(pid) = claimed { + Self::release_send_quota(&flow, pid).await; } + tracing::debug!( + packet_id = ?packet_id, + "Connection changed before publish was sent; conforming it to the current connection" + ); + return self + .conform_to_connection(ready.request) + .await + .map(Transmitted::Restaged); } - let ack = self.setup_publish_acknowledgment(qos, packet_id); + let ReadyPublish { + packet: publish, + reservation, + .. + } = *ready; + let qos = publish.qos; + + let in_flight = match packet_id { + Some(pid) if qos != QoS::AtMostOnce => { + let stored = match self.with_aliased_topic(publish.clone()).await { + Ok(stored) => { + self.session + .read() + .await + .store_unacked_publish(stored) + .await + } + Err(e) => Err(e), + }; + if let Err(e) = stored { + Self::release_send_quota(&flow, pid).await; + return Err(e); + } + let (completion, handle) = Completion::new(); + self.publish_outcomes.lock().track(pid, qos, completion); + Some(handle) + } + _ => None, + }; + drop(reservation); if publish.payload.len() > 10000 { tracing::debug!( @@ -998,11 +1194,27 @@ impl DirectClientInner { .topic_alias() .filter(|_| !publish.topic_name.is_empty()) .map(|alias| (alias, publish.topic_name.clone())); - self.send_publish_packet(publish).await?; - if let Some((alias, alias_topic)) = alias_mapping { + let link = self + .connection_alive + .as_ref() + .map_or_else(|| watch::channel(false).1, |alive| alive.subscribe()); + let written = self.send_publish_packet(publish).await; + if let Some((alias, alias_topic)) = alias_mapping.filter(|_| written.is_ok()) { self.record_outbound_topic_alias(alias, &alias_topic).await; } - Ok(ack) + match (written, in_flight) { + (Ok(()), None) => Ok(Transmitted::Sent(Delivery::of(qos, packet_id))), + (Ok(()), Some(handle)) => Ok(Transmitted::InFlight(InFlight { handle, link })), + (Err(e), None) => Err(e), + (Err(e), Some(handle)) => { + tracing::debug!( + packet_id = ?packet_id, + error = %e, + "Stored PUBLISH could not be written; it is re-sent with the session" + ); + Ok(Transmitted::Detached(handle)) + } + } } async fn record_outbound_topic_alias(&self, alias: u16, topic: &str) { @@ -1247,12 +1459,14 @@ impl DirectClientInner { let writer = self.writer.as_ref().ok_or(MqttError::NotConnected)?; - let packet_id = self.allocate_packet_id().await?; + let reservation = self.allocate_packet_id().await?; + let packet_id = reservation.packet_id(); let mut packet = packet; packet.packet_id = packet_id; let (tx, rx) = oneshot::channel(); self.pending_subacks.lock().insert(packet_id, tx); + drop(reservation); maybe_store_subscriptions( &self.stored_subscriptions, @@ -1325,12 +1539,14 @@ impl DirectClientInner { let writer = self.writer.as_ref().ok_or(MqttError::NotConnected)?; - let packet_id = self.allocate_packet_id().await?; + let reservation = self.allocate_packet_id().await?; + let packet_id = reservation.packet_id(); let mut packet = packet; packet.packet_id = packet_id; let (tx, rx) = oneshot::channel(); self.pending_unsubacks.lock().insert(packet_id, tx); + drop(reservation); { let mut stored = self.stored_subscriptions.lock(); @@ -1462,6 +1678,26 @@ impl DirectClientInner { will.map_or_else(Properties::default, |w| w.properties.clone().into()) } + fn connection_lifecycle( + &self, + connection_epoch: u64, + reader_task: Arc>>, + ) -> Result { + Ok(keepalive::ConnectionLifecycle { + connected: self.connected.clone(), + connection_epoch, + current_connection_epoch: self.connection_epoch.clone(), + callbacks: Arc::clone(&self.connection_event_callbacks), + closing: Arc::new(AtomicBool::new(false)), + alive: self + .connection_alive + .clone() + .ok_or(MqttError::NotConnected)?, + flow: Arc::clone(&self.send_flow), + reader_task, + }) + } + fn start_background_tasks( &mut self, reader: UnifiedReader, @@ -1471,15 +1707,10 @@ impl DirectClientInner { let reader_callbacks = self.callback_manager.clone(); let suback_channels = self.pending_subacks.clone(); let unsuback_channels = self.pending_unsubacks.clone(); - let puback_channels = self.pending_pubacks.clone(); - let pubcomp_channels = self.pending_pubcomps.clone(); + let publish_outcomes = Arc::clone(&self.publish_outcomes); + let reader_task = Arc::new(Mutex::new(None)); let writer_for_keepalive = self.writer.as_ref().ok_or(MqttError::NotConnected)?.clone(); - let lifecycle = keepalive::ConnectionLifecycle { - connected: self.connected.clone(), - connection_epoch, - current_connection_epoch: self.connection_epoch.clone(), - callbacks: Arc::clone(&self.connection_event_callbacks), - }; + let lifecycle = self.connection_lifecycle(connection_epoch, Arc::clone(&reader_task))?; let writer_for_reader = writer_for_keepalive.clone(); let keepalive_state = self.keepalive_state.clone(); @@ -1489,12 +1720,18 @@ impl DirectClientInner { callback_manager: reader_callbacks, suback_channels, unsuback_channels, - puback_channels, - pubcomp_channels, + publish_outcomes, writer: writer_for_reader, lifecycle: lifecycle.clone(), #[cfg(feature = "transport-quic")] protocol_version: self.options.protocol_version.as_u8(), + #[cfg(feature = "transport-quic")] + maximum_packet_size: self + .options + .properties + .maximum_packet_size + .and_then(|size| usize::try_from(size).ok()) + .unwrap_or(usize::MAX), auth_handler: self.auth_handler.clone(), auth_method: self.auth_method.clone(), keepalive_state: keepalive_state.clone(), @@ -1513,11 +1750,13 @@ impl DirectClientInner { }; let ctx_for_packet_reader = ctx.clone(); - self.packet_reader_handle = Some(tokio::spawn(async move { + let packet_reader = tokio::spawn(async move { tracing::debug!("📦 PACKET READER - Task starting"); packet_reader_task_with_responses(reader, ctx_for_packet_reader).await; tracing::debug!("📦 PACKET READER - Task exited"); - })); + }); + *reader_task.lock() = Some(packet_reader.abort_handle()); + self.packet_reader_handle = Some(packet_reader); let keepalive_interval = self.negotiated_keep_alive(); if keepalive_interval.is_zero() { @@ -1836,17 +2075,17 @@ pub mod tests { }, ) .await; + let Ok(StagedPublish::Queued(handle)) = within else { + panic!("within-limit publish must queue: {within:?}"); + }; assert!( - matches!( - within, - Ok(StagedPublish::Queued(PublishResult::QoS1Or2 { .. })) - ), - "within-limit publish must queue: {within:?}" + !client.queued_messages.lock().is_empty(), + "within-limit publish must be queued" ); assert_eq!( - client.queued_messages.lock().len(), - 1, - "within-limit publish must be queued" + handle.try_outcome(), + None, + "a queued publish must not report an outcome before it is flushed" ); } @@ -1978,4 +2217,965 @@ pub mod tests { let session = client.session.write().await; assert_eq!(session.client_id(), "test-client"); } + + type SeenFrames = Arc>>; + + async fn read_frame(stream: &mut tokio::net::TcpStream) -> Option> { + use tokio::io::AsyncReadExt; + let mut first = [0u8; 1]; + stream.read_exact(&mut first).await.ok()?; + let mut len = 0usize; + let mut shift = 0; + loop { + let mut b = [0u8; 1]; + stream.read_exact(&mut b).await.ok()?; + len |= usize::from(b[0] & 0x7f) << shift; + shift += 7; + if b[0] & 0x80 == 0 { + break; + } + } + let mut body = vec![0u8; len]; + stream.read_exact(&mut body).await.ok()?; + let mut out = vec![first[0]]; + out.extend(body); + Some(out) + } + + async fn silent_broker(connacks: Vec>) -> (std::net::SocketAddr, SeenFrames) { + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let seen: SeenFrames = Arc::new(Mutex::new(Vec::new())); + let seen_broker = Arc::clone(&seen); + tokio::spawn(async move { + for (conn, connack) in connacks.into_iter().enumerate() { + let (mut s, _) = listener.accept().await.unwrap(); + read_frame(&mut s).await.unwrap(); + s.write_all(&connack).await.unwrap(); + let seen_conn = Arc::clone(&seen_broker); + tokio::spawn(async move { + while let Some(f) = read_frame(&mut s).await { + seen_conn.lock().push((conn, f[0])); + } + }); + } + }); + (addr, seen) + } + + async fn connect_to(client: &mut DirectClientInner, addr: std::net::SocketAddr) -> bool { + use mqtt5_protocol::Transport; + let mut transport = crate::transport::tcp::TcpTransport::from_addr(addr); + transport.connect().await.unwrap(); + client + .connect(TransportType::Tcp(transport)) + .await + .unwrap() + .session_present + } + + fn publishes_on(seen: &SeenFrames, conn: usize) -> Vec { + seen.lock() + .iter() + .filter(|(c, first)| *c == conn && first >> 4 == 3) + .map(|(_, first)| *first) + .collect() + } + + fn qos(level: QoS) -> PublishOptions { + PublishOptions { + qos: level, + ..Default::default() + } + } + + async fn claim_quota(client: &DirectClientInner, ready: &ReadyPublish) -> Option { + let flow = Arc::clone(client.session.read().await.flow_control()); + match ready.packet_id() { + Some(packet_id) => Some( + FlowControlManager::acquire_shared_send_quota(&flow, packet_id) + .await + .unwrap(), + ), + None => None, + } + } + + async fn send_with_claim( + client: &DirectClientInner, + mut ready: Box, + mut claim: Option, + ) -> usize { + let mut restages = 0; + loop { + match client.transmit_publish(ready, claim).await.unwrap() { + Transmitted::Sent(_) | Transmitted::InFlight(_) | Transmitted::Detached(_) => { + return restages; + } + Transmitted::Restaged(next) => { + restages += 1; + claim = claim_quota(client, &next).await; + ready = next; + } + } + } + } + + #[tokio::test] + async fn stale_quota_claim_does_not_exceed_receive_maximum_after_reconnect() { + let receive_maximum_one = + |session_present: u8| vec![0x20, 0x06, session_present, 0x00, 0x03, 0x21, 0x00, 0x01]; + let (addr, seen) = + silent_broker(vec![receive_maximum_one(0), receive_maximum_one(1)]).await; + let mut client = + DirectClientInner::new(ConnectOptions::new("stale-quota").with_clean_start(false)); + connect_to(&mut client, addr).await; + + let Ok(StagedPublish::Ready(b)) = client + .stage_publish("t/b".into(), b"b".to_vec(), qos(QoS::AtLeastOnce)) + .await + else { + panic!("publish b must stage while connected"); + }; + let flow = Arc::clone(client.session.read().await.flow_control()); + let stale_claim = claim_quota(&client, &b).await; + + assert!(connect_to(&mut client, addr).await); + + assert_eq!(send_with_claim(&client, b, stale_claim).await, 1); + assert_eq!(flow.read().await.in_flight_count().await, 1); + + let Ok(StagedPublish::Ready(c)) = client + .stage_publish("t/c".into(), b"c".to_vec(), qos(QoS::AtLeastOnce)) + .await + else { + panic!("publish c must stage while connected"); + }; + let c_quota = tokio::time::timeout( + Duration::from_millis(500), + FlowControlManager::acquire_shared_send_quota(&flow, c.packet_id().unwrap()), + ) + .await; + if let Ok(Ok(generation)) = c_quota { + send_with_claim(&client, c, Some(generation)).await; + } + tokio::time::sleep(Duration::from_millis(200)).await; + let conn2 = publishes_on(&seen, 1); + assert!( + conn2.len() <= 1, + "client exceeded server Receive Maximum 1 on the new connection: {conn2:x?}" + ); + } + + #[tokio::test] + async fn publish_waiting_across_reconnect_conforms_to_new_connection() { + let first = vec![0x20, 0x06, 0x00, 0x00, 0x03, 0x21, 0x00, 0x01]; + let maximum_qos_one = vec![0x20, 0x05, 0x01, 0x00, 0x02, 0x24, 0x01]; + let (addr, seen) = silent_broker(vec![first, maximum_qos_one]).await; + let mut client = + DirectClientInner::new(ConnectOptions::new("reconform").with_clean_start(false)); + connect_to(&mut client, addr).await; + let flow = Arc::clone(client.session.read().await.flow_control()); + + let Ok(StagedPublish::Ready(a)) = client + .stage_publish("t/a".into(), b"a".to_vec(), qos(QoS::AtLeastOnce)) + .await + else { + panic!("publish a must stage while connected"); + }; + let a_claim = claim_quota(&client, &a).await; + assert_eq!(send_with_claim(&client, a, a_claim).await, 0); + + let Ok(StagedPublish::Ready(b)) = client + .stage_publish("t/b".into(), b"b".to_vec(), qos(QoS::ExactlyOnce)) + .await + else { + panic!("publish b must stage while connected"); + }; + let b_id = b.packet_id().unwrap(); + let waiting_flow = Arc::clone(&flow); + let waiting = tokio::spawn(async move { + FlowControlManager::acquire_shared_send_quota(&waiting_flow, b_id).await + }); + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(!waiting.is_finished(), "b must wait for Receive Maximum 1"); + + assert!(connect_to(&mut client, addr).await); + let b_claim = waiting.await.unwrap().unwrap(); + assert_eq!(send_with_claim(&client, b, Some(b_claim)).await, 1); + + tokio::time::sleep(Duration::from_millis(200)).await; + assert_eq!( + publishes_on(&seen, 1), + vec![0x3A, 0x32], + "the replayed QoS 1 publish goes first, then b downgraded to the new Maximum QoS 1" + ); + } + + async fn unacked_packet_ids(client: &DirectClientInner) -> Vec { + client + .session + .read() + .await + .get_unacked_publishes() + .await + .iter() + .filter_map(|publish| publish.packet_id) + .collect() + } + + #[tokio::test] + async fn stale_flush_task_does_not_write_after_reconnect() { + let resume_receive_maximum_one = vec![0x20, 0x06, 0x01, 0x00, 0x03, 0x21, 0x00, 0x01]; + let resume = vec![0x20, 0x03, 0x01, 0x00, 0x00]; + let (addr, seen) = silent_broker(vec![ + vec![0x20, 0x03, 0x00, 0x00, 0x00], + resume_receive_maximum_one, + resume, + ]) + .await; + let mut client = + DirectClientInner::new(ConnectOptions::new("stale-flush").with_clean_start(false)); + connect_to(&mut client, addr).await; + client.set_connected(false); + for topic in ["t/a", "t/b"] { + let queued = client + .stage_publish(topic.into(), b"x".to_vec(), qos(QoS::AtLeastOnce)) + .await; + assert!(matches!(queued, Ok(StagedPublish::Queued(_)))); + } + + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(100)).await; + assert_eq!( + publishes_on(&seen, 1).len(), + 1, + "setup: a flushed, b waits for quota" + ); + let stale_writer = Arc::clone(client.writer.as_ref().unwrap()); + let held = stale_writer.lock().await; + let flow = Arc::clone(client.session.read().await.flow_control()); + let a = unacked_packet_ids(&client).await[0]; + flow.read().await.acknowledge(a).await.unwrap(); + tokio::time::sleep(Duration::from_millis(100)).await; + + assert!(connect_to(&mut client, addr).await); + drop(held); + tokio::time::sleep(Duration::from_millis(200)).await; + + assert_eq!( + publishes_on(&seen, 1).len(), + 1, + "the flush task of the replaced connection wrote after the reconnect" + ); + } + + #[tokio::test] + async fn abandoned_qos2_replay_id_is_quarantined_until_session_is_lost() { + let first = vec![0x20, 0x03, 0x00, 0x00, 0x00]; + let resume = vec![0x20, 0x03, 0x01, 0x00, 0x00]; + let resume_maximum_qos_one = vec![0x20, 0x05, 0x01, 0x00, 0x02, 0x24, 0x01]; + let lost = vec![0x20, 0x03, 0x00, 0x00, 0x00]; + let (addr, seen) = silent_broker(vec![first, resume, resume_maximum_qos_one, lost]).await; + let mut client = + DirectClientInner::new(ConnectOptions::new("quarantine").with_clean_start(false)); + connect_to(&mut client, addr).await; + client.set_connected(false); + let queued = client + .stage_publish("t/two".into(), b"x".to_vec(), qos(QoS::ExactlyOnce)) + .await; + assert!(matches!(queued, Ok(StagedPublish::Queued(_)))); + + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(100)).await; + let abandoned = unacked_packet_ids(&client).await[0]; + assert_eq!(publishes_on(&seen, 1), vec![0x34]); + + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(100)).await; + assert!( + publishes_on(&seen, 2).is_empty(), + "QoS 2 replay after Maximum QoS 1" + ); + assert!(unacked_packet_ids(&client).await.is_empty()); + + let mut handed_out = false; + for _ in 0..u16::MAX { + let reservation = client.allocate_packet_id().await.unwrap(); + handed_out |= reservation.packet_id() == abandoned; + } + assert!( + !handed_out, + "a quarantined QoS 2 identifier was reallocated" + ); + + assert!(!connect_to(&mut client, addr).await); + let mut released = false; + for _ in 0..u16::MAX { + let reservation = client.allocate_packet_id().await.unwrap(); + released |= reservation.packet_id() == abandoned; + } + assert!(released, "Session Present 0 must release the quarantine"); + } + + #[derive(Clone, Copy)] + enum AckScript { + Acknowledged, + ReceiptRefused, + Completed, + } + + async fn ack_processed_while_disconnecting(script: AckScript) -> Option { + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let (pid_tx, mut pid_rx) = tokio::sync::mpsc::unbounded_channel::(); + let (ack_tx, ack_rx) = oneshot::channel::<()>(); + tokio::spawn(async move { + let (mut s0, _) = listener.accept().await.unwrap(); + read_frame(&mut s0).await.unwrap(); + s0.write_all(&[0x20, 0x03, 0x00, 0x00, 0x00]).await.unwrap(); + tokio::spawn(async move { while read_frame(&mut s0).await.is_some() {} }); + let (mut s1, _) = listener.accept().await.unwrap(); + read_frame(&mut s1).await.unwrap(); + s1.write_all(&[0x20, 0x03, 0x01, 0x00, 0x00]).await.unwrap(); + let pid = loop { + let f = read_frame(&mut s1).await.unwrap(); + if f[0] >> 4 == 3 { + let tlen = usize::from(u16::from_be_bytes([f[1], f[2]])); + break u16::from_be_bytes([f[3 + tlen], f[4 + tlen]]); + } + }; + let [hi, lo] = pid.to_be_bytes(); + if matches!(script, AckScript::Completed) { + s1.write_all(&[0x50, 0x02, hi, lo]).await.unwrap(); + while read_frame(&mut s1).await.unwrap()[0] != 0x62 {} + } + pid_tx.send(pid).unwrap(); + ack_rx.await.unwrap(); + let ack: &[u8] = match script { + AckScript::Acknowledged => &[0x40, 0x02, hi, lo], + AckScript::ReceiptRefused => &[0x50, 0x03, hi, lo, 0x80], + AckScript::Completed => &[0x70, 0x02, hi, lo], + }; + s1.write_all(ack).await.unwrap(); + tokio::spawn(async move { while read_frame(&mut s1).await.is_some() {} }); + let (mut s2, _) = listener.accept().await.unwrap(); + read_frame(&mut s2).await.unwrap(); + s2.write_all(&[0x20, 0x03, 0x01, 0x00, 0x00]).await.unwrap(); + while read_frame(&mut s2).await.is_some() {} + }); + + let level = match script { + AckScript::Acknowledged => QoS::AtLeastOnce, + AckScript::ReceiptRefused | AckScript::Completed => QoS::ExactlyOnce, + }; + let mut client = + DirectClientInner::new(ConnectOptions::new("ack-abort").with_clean_start(false)); + connect_to(&mut client, addr).await; + client.set_connected(false); + let Ok(StagedPublish::Queued(handle)) = client + .stage_publish("t/a".into(), b"a".to_vec(), qos(level)) + .await + else { + panic!("setup: offline publish must queue"); + }; + assert!(connect_to(&mut client, addr).await); + pid_rx.recv().await.unwrap(); + + let flow = Arc::clone(client.session.read().await.flow_control()); + let guard = Arc::clone(&flow).write_owned().await; + ack_tx.send(()).unwrap(); + tokio::time::sleep(Duration::from_millis(200)).await; + tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(500)).await; + drop(guard); + }); + client.disconnect().await.unwrap(); + + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(300)).await; + handle.try_outcome() + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn puback_processed_while_disconnecting_settles_the_handle() { + assert!(matches!( + ack_processed_while_disconnecting(AckScript::Acknowledged).await, + Some(PublishOutcome::Delivered(Delivery::AtLeastOnce { .. })) + )); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn refused_pubrec_processed_while_disconnecting_settles_the_handle() { + assert_eq!( + ack_processed_while_disconnecting(AckScript::ReceiptRefused).await, + Some(PublishOutcome::Rejected(PublishRejection::Refused( + ReasonCode::UnspecifiedError + ))) + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn pubcomp_processed_while_disconnecting_settles_the_handle() { + assert!(matches!( + ack_processed_while_disconnecting(AckScript::Completed).await, + Some(PublishOutcome::Delivered(Delivery::ExactlyOnce { .. })) + )); + } + + #[tokio::test] + async fn publish_staged_before_reconnect_is_restaged_even_with_a_current_quota_claim() { + let first = vec![0x20, 0x03, 0x00, 0x00, 0x00]; + let maximum_qos_one = vec![0x20, 0x05, 0x01, 0x00, 0x02, 0x24, 0x01]; + let (addr, seen) = silent_broker(vec![first, maximum_qos_one]).await; + let mut client = + DirectClientInner::new(ConnectOptions::new("epoch-only").with_clean_start(false)); + connect_to(&mut client, addr).await; + let Ok(StagedPublish::Ready(staged)) = client + .stage_publish("t/e".into(), b"e".to_vec(), qos(QoS::ExactlyOnce)) + .await + else { + panic!("publish must stage while connected"); + }; + + assert!(connect_to(&mut client, addr).await); + let current_claim = claim_quota(&client, &staged).await; + assert_eq!(send_with_claim(&client, staged, current_claim).await, 1); + + tokio::time::sleep(Duration::from_millis(200)).await; + assert_eq!( + publishes_on(&seen, 1), + vec![0x32], + "a publish staged on the previous connection must be conformed to the new Maximum QoS 1" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn abandoned_qos2_id_is_quarantined_before_it_leaves_the_session_store() { + let first = vec![0x20, 0x03, 0x00, 0x00, 0x00]; + let resume = vec![0x20, 0x03, 0x01, 0x00, 0x00]; + let resume_maximum_qos_one = vec![0x20, 0x05, 0x01, 0x00, 0x02, 0x24, 0x01]; + let (addr, _seen) = silent_broker(vec![first, resume, resume_maximum_qos_one]).await; + let mut client = + DirectClientInner::new(ConnectOptions::new("quarantine-order").with_clean_start(false)); + connect_to(&mut client, addr).await; + client.set_connected(false); + let queued = client + .stage_publish("t/two".into(), b"x".to_vec(), qos(QoS::ExactlyOnce)) + .await; + assert!(matches!(queued, Ok(StagedPublish::Queued(_)))); + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(100)).await; + let abandoned = unacked_packet_ids(&client).await[0]; + + let ids = Arc::clone(&client.outbound_ids); + let (locked_tx, locked_rx) = oneshot::channel::<()>(); + let (inspect_tx, inspect_rx) = std::sync::mpsc::channel::<()>(); + let (state_tx, state_rx) = oneshot::channel::(); + let holder = std::thread::spawn(move || { + let held = ids.lock(); + let _ = locked_tx.send(()); + let _ = inspect_rx.recv(); + let _ = state_tx.send(held.holds(abandoned)); + }); + locked_rx.await.unwrap(); + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(300)).await; + let in_store = unacked_packet_ids(&client).await.contains(&abandoned); + inspect_tx.send(()).unwrap(); + let quarantined = state_rx.await.unwrap(); + holder.join().unwrap(); + + assert!( + in_store || quarantined, + "packet id {abandoned} left the session store before it was quarantined; an allocation in that window could reuse it" + ); + tokio::time::sleep(Duration::from_millis(100)).await; + assert!(client.outbound_ids.lock().holds(abandoned)); + } + + #[tokio::test] + async fn downgraded_message_is_requeued_when_the_connection_ends_before_it_is_written() { + let first = vec![0x20, 0x03, 0x00, 0x00, 0x00]; + let resume_maximum_qos_zero = vec![0x20, 0x05, 0x01, 0x00, 0x02, 0x24, 0x00]; + let resume = vec![0x20, 0x03, 0x01, 0x00, 0x00]; + let (addr, seen) = silent_broker(vec![first, resume_maximum_qos_zero, resume]).await; + let mut client = DirectClientInner::new( + ConnectOptions::new("downgrade-requeue").with_clean_start(false), + ); + connect_to(&mut client, addr).await; + client.set_connected(false); + let Ok(StagedPublish::Queued(handle)) = client + .stage_publish("t/one".into(), b"x".to_vec(), qos(QoS::AtLeastOnce)) + .await + else { + panic!("setup: offline publish must queue"); + }; + + assert!(connect_to(&mut client, addr).await); + let replaced_writer = Arc::clone(client.writer.as_ref().unwrap()); + let held = replaced_writer.lock().await; + tokio::time::sleep(Duration::from_millis(100)).await; + assert!(connect_to(&mut client, addr).await); + drop(held); + tokio::time::sleep(Duration::from_millis(300)).await; + + assert!( + publishes_on(&seen, 1).is_empty(), + "nothing reaches the replaced connection" + ); + assert_eq!( + publishes_on(&seen, 2), + vec![0x32], + "the message that never reached the wire is re-queued and sent on the next connection" + ); + assert_eq!(handle.try_outcome(), None); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn stuck_replay_on_replaced_connection_does_not_starve_the_new_one() { + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let seen: SeenFrames = Arc::new(Mutex::new(Vec::new())); + let seen_broker = Arc::clone(&seen); + let (hold_tx, mut hold_rx) = tokio::sync::mpsc::unbounded_channel(); + tokio::spawn(async move { + let connacks = [ + vec![0x20, 0x03, 0x00, 0x00, 0x00], + vec![0x20, 0x03, 0x01, 0x00, 0x00], + vec![0x20, 0x03, 0x01, 0x00, 0x00], + ]; + for (conn, connack) in connacks.into_iter().enumerate() { + let (mut s, _) = listener.accept().await.unwrap(); + read_frame(&mut s).await.unwrap(); + s.write_all(&connack).await.unwrap(); + if conn == 1 { + let _ = hold_tx.send(s); + continue; + } + let seen_conn = Arc::clone(&seen_broker); + tokio::spawn(async move { + while let Some(f) = read_frame(&mut s).await { + seen_conn.lock().push((conn, f[0])); + } + }); + } + }); + let mut client = + DirectClientInner::new(ConnectOptions::new("stuck-replay").with_clean_start(false)); + connect_to(&mut client, addr).await; + client.set_connected(false); + for i in 0..64 { + let queued = client + .stage_publish( + format!("t/{i}"), + vec![0u8; 256 * 1024], + qos(QoS::AtLeastOnce), + ) + .await; + assert!(matches!(queued, Ok(StagedPublish::Queued(_)))); + } + assert!(connect_to(&mut client, addr).await); + let _held = hold_rx.recv().await.unwrap(); + tokio::time::sleep(Duration::from_millis(500)).await; + + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(1500)).await; + let on_new = publishes_on(&seen, 2).len(); + let Ok(StagedPublish::Ready(live)) = client + .stage_publish("t/live".into(), b"x".to_vec(), qos(QoS::AtLeastOnce)) + .await + else { + panic!("live publish must stage"); + }; + let flow = Arc::clone(client.session.read().await.flow_control()); + let quota = tokio::time::timeout( + Duration::from_secs(3), + FlowControlManager::acquire_shared_send_quota(&flow, live.packet_id().unwrap()), + ) + .await; + assert!( + on_new > 0 && quota.is_ok(), + "new healthy connection is starved by a replay stuck on the replaced connection: resent on new={on_new}, live quota={quota:?}" + ); + } + + #[tokio::test] + async fn stored_publish_whose_write_fails_is_returned_detached() { + let (addr, _seen) = silent_broker(vec![vec![0x20, 0x03, 0x00, 0x00, 0x00]]).await; + let mut client = + DirectClientInner::new(ConnectOptions::new("detached").with_clean_start(false)); + connect_to(&mut client, addr).await; + let Ok(StagedPublish::Ready(ready)) = client + .stage_publish("t/d".into(), b"d".to_vec(), qos(QoS::AtLeastOnce)) + .await + else { + panic!("publish must stage while connected"); + }; + let packet_id = ready.packet_id().unwrap(); + let claim = claim_quota(&client, &ready).await; + client + .writer + .as_ref() + .unwrap() + .lock() + .await + .close(None) + .await + .unwrap(); + + let transmitted = client.transmit_publish(ready, claim).await; + let Ok(Transmitted::Detached(handle)) = transmitted else { + panic!("a stored publish whose write failed must come back as a pending handle"); + }; + assert_eq!(handle.try_outcome(), None); + assert_eq!(unacked_packet_ids(&client).await, vec![packet_id]); + } + + #[tokio::test] + async fn disconnect_while_already_disconnected_closes_send_quota() { + let receive_maximum_one = vec![0x20, 0x06, 0x00, 0x00, 0x03, 0x21, 0x00, 0x01]; + let (addr, _seen) = silent_broker(vec![receive_maximum_one]).await; + let mut client = + DirectClientInner::new(ConnectOptions::new("disconnect-quota").with_clean_start(false)); + connect_to(&mut client, addr).await; + let Ok(StagedPublish::Ready(first)) = client + .stage_publish("t/1".into(), b"1".to_vec(), qos(QoS::AtLeastOnce)) + .await + else { + panic!("publish must stage while connected"); + }; + let claim = claim_quota(&client, &first).await; + send_with_claim(&client, first, claim).await; + let flow = Arc::clone(client.session.read().await.flow_control()); + let waiting = tokio::spawn(async move { + FlowControlManager::acquire_shared_send_quota(&flow, 4242).await + }); + tokio::time::sleep(Duration::from_millis(50)).await; + assert!(!waiting.is_finished(), "setup: the send quota is exhausted"); + + client.set_connected(false); + assert!(matches!( + client.disconnect().await, + Err(MqttError::NotConnected) + )); + let waited = tokio::time::timeout(Duration::from_secs(2), waiting).await; + assert!( + matches!(waited, Ok(Ok(Err(MqttError::NotConnected)))), + "disconnect must release publishes waiting for send quota: {waited:?}" + ); + } + + async fn mismatched_acknowledgement(level: QoS, ack_type: u8) -> (Vec, bool) { + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let (frames_tx, frames_rx) = oneshot::channel::>>(); + tokio::spawn(async move { + let (mut s, _) = listener.accept().await.unwrap(); + read_frame(&mut s).await.unwrap(); + s.write_all(&[0x20, 0x03, 0x00, 0x00, 0x00]).await.unwrap(); + let pid = loop { + let f = read_frame(&mut s).await.unwrap(); + if f[0] >> 4 == 3 { + let tlen = usize::from(u16::from_be_bytes([f[1], f[2]])); + break u16::from_be_bytes([f[3 + tlen], f[4 + tlen]]); + } + }; + let [hi, lo] = pid.to_be_bytes(); + s.write_all(&[ack_type, 0x02, hi, lo]).await.unwrap(); + let mut after = Vec::new(); + while let Ok(Some(f)) = + tokio::time::timeout(Duration::from_millis(500), read_frame(&mut s)).await + { + after.push(f); + } + let _ = frames_tx.send(after); + }); + let mut client = + DirectClientInner::new(ConnectOptions::new("ack-mismatch").with_clean_start(false)); + connect_to(&mut client, addr).await; + let Ok(StagedPublish::Ready(ready)) = client + .stage_publish("t/m".into(), b"m".to_vec(), qos(level)) + .await + else { + panic!("publish must stage while connected"); + }; + let packet_id = ready.packet_id().unwrap(); + let claim = claim_quota(&client, &ready).await; + send_with_claim(&client, ready, claim).await; + let after = frames_rx.await.unwrap(); + let disconnect = after + .into_iter() + .find(|f| f[0] == 0xE0) + .map(|f| f[1..].to_vec()) + .unwrap_or_default(); + let still_held = unacked_packet_ids(&client).await.contains(&packet_id); + (disconnect, still_held) + } + + #[tokio::test] + async fn puback_for_qos2_publish_is_a_protocol_error() { + let (disconnect, still_held) = mismatched_acknowledgement(QoS::ExactlyOnce, 0x40).await; + assert_eq!(disconnect.first(), Some(&0x82)); + assert!( + still_held, + "a mismatched PUBACK must not release the QoS 2 state" + ); + } + + #[tokio::test] + async fn pubcomp_for_qos1_publish_is_a_protocol_error() { + let (disconnect, still_held) = mismatched_acknowledgement(QoS::AtLeastOnce, 0x70).await; + assert_eq!(disconnect.first(), Some(&0x82)); + assert!( + still_held, + "a mismatched PUBCOMP must not release the QoS 1 state" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn refused_pubrec_aborted_before_removal_keeps_the_publish_unsettled() { + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let (pid_tx, mut pid_rx) = tokio::sync::mpsc::unbounded_channel::(); + let (ack_tx, ack_rx) = oneshot::channel::<()>(); + let later: SeenFrames = Arc::new(Mutex::new(Vec::new())); + let later_broker = Arc::clone(&later); + tokio::spawn(async move { + let (mut s1, _) = listener.accept().await.unwrap(); + read_frame(&mut s1).await.unwrap(); + s1.write_all(&[0x20, 0x03, 0x00, 0x00, 0x00]).await.unwrap(); + let pid = loop { + let f = read_frame(&mut s1).await.unwrap(); + if f[0] >> 4 == 3 { + let tlen = usize::from(u16::from_be_bytes([f[1], f[2]])); + break u16::from_be_bytes([f[3 + tlen], f[4 + tlen]]); + } + }; + let [hi, lo] = pid.to_be_bytes(); + pid_tx.send(pid).unwrap(); + ack_rx.await.unwrap(); + s1.write_all(&[0x50, 0x03, hi, lo, 0x80]).await.unwrap(); + tokio::spawn(async move { while read_frame(&mut s1).await.is_some() {} }); + let (mut s2, _) = listener.accept().await.unwrap(); + read_frame(&mut s2).await.unwrap(); + s2.write_all(&[0x20, 0x03, 0x01, 0x00, 0x00]).await.unwrap(); + while let Some(f) = read_frame(&mut s2).await { + later_broker.lock().push((1, f[0])); + } + }); + let mut client = + DirectClientInner::new(ConnectOptions::new("pubrec-abort").with_clean_start(false)); + connect_to(&mut client, addr).await; + let Ok(StagedPublish::Ready(ready)) = client + .stage_publish("t/a".into(), b"a".to_vec(), qos(QoS::ExactlyOnce)) + .await + else { + panic!("publish must stage"); + }; + let claim = claim_quota(&client, &ready).await; + let Ok(Transmitted::InFlight(in_flight)) = client.transmit_publish(ready, claim).await + else { + panic!("live publish must be in flight"); + }; + let handle = in_flight.handle.clone(); + pid_rx.recv().await.unwrap(); + let session = Arc::clone(&client.session); + let guard = session.read_owned().await; + ack_tx.send(()).unwrap(); + tokio::time::sleep(Duration::from_millis(200)).await; + let outcome_at_abort = handle.try_outcome(); + tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(500)).await; + drop(guard); + }); + client.disconnect().await.unwrap(); + assert!(connect_to(&mut client, addr).await); + tokio::time::sleep(Duration::from_millis(300)).await; + let resent = publishes_on(&later, 1); + assert_eq!( + outcome_at_abort, None, + "a PUBREC refusal that was not applied to the session must not settle the publish" + ); + assert!( + !resent.is_empty(), + "the still-stored PUBLISH is replayed on resume" + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn stuck_downgraded_write_does_not_block_reconnect() { + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let seen: SeenFrames = Arc::new(Mutex::new(Vec::new())); + let seen_broker = Arc::clone(&seen); + let (hold_tx, mut hold_rx) = tokio::sync::mpsc::unbounded_channel(); + tokio::spawn(async move { + let connacks = [ + vec![0x20, 0x03, 0x00, 0x00, 0x00], + vec![0x20, 0x05, 0x01, 0x00, 0x02, 0x24, 0x00], + vec![0x20, 0x03, 0x01, 0x00, 0x00], + ]; + for (conn, connack) in connacks.into_iter().enumerate() { + let (mut s, _) = listener.accept().await.unwrap(); + read_frame(&mut s).await.unwrap(); + s.write_all(&connack).await.unwrap(); + if conn == 1 { + let _ = hold_tx.send(s); + continue; + } + let seen_conn = Arc::clone(&seen_broker); + tokio::spawn(async move { + while let Some(f) = read_frame(&mut s).await { + seen_conn.lock().push((conn, f[0])); + } + }); + } + }); + let mut client = + DirectClientInner::new(ConnectOptions::new("stuck-downgrade").with_clean_start(false)); + connect_to(&mut client, addr).await; + client.set_connected(false); + let mut handles = Vec::new(); + for i in 0..64 { + let Ok(StagedPublish::Queued(h)) = client + .stage_publish( + format!("t/{i}"), + vec![0u8; 256 * 1024], + qos(QoS::AtLeastOnce), + ) + .await + else { + panic!("queue"); + }; + handles.push(h); + } + assert!(connect_to(&mut client, addr).await); + let _held = hold_rx.recv().await.unwrap(); + tokio::time::sleep(Duration::from_millis(500)).await; + let reconnected = + tokio::time::timeout(Duration::from_secs(3), connect_to(&mut client, addr)).await; + assert!( + reconnected.is_ok(), + "reconnect blocked behind a stuck downgraded write" + ); + tokio::time::sleep(Duration::from_millis(1500)).await; + let on_new = publishes_on(&seen, 2).len(); + assert!( + on_new > 0, + "remaining queued messages flushed on the new connection" + ); + } + + #[tokio::test] + async fn refused_pubrec_after_pubrel_is_a_protocol_error() { + use tokio::io::AsyncWriteExt; + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let (frames_tx, frames_rx) = oneshot::channel::>>(); + tokio::spawn(async move { + let (mut s, _) = listener.accept().await.unwrap(); + read_frame(&mut s).await.unwrap(); + s.write_all(&[0x20, 0x03, 0x00, 0x00, 0x00]).await.unwrap(); + let pid = loop { + let f = read_frame(&mut s).await.unwrap(); + if f[0] >> 4 == 3 { + let tlen = usize::from(u16::from_be_bytes([f[1], f[2]])); + break u16::from_be_bytes([f[3 + tlen], f[4 + tlen]]); + } + }; + let [hi, lo] = pid.to_be_bytes(); + s.write_all(&[0x50, 0x02, hi, lo]).await.unwrap(); + while read_frame(&mut s).await.unwrap()[0] != 0x62 {} + s.write_all(&[0x50, 0x03, hi, lo, 0x80]).await.unwrap(); + let mut after = Vec::new(); + while let Ok(Some(f)) = + tokio::time::timeout(Duration::from_millis(500), read_frame(&mut s)).await + { + after.push(f); + } + let _ = frames_tx.send(after); + }); + let mut client = DirectClientInner::new( + ConnectOptions::new("pubrec-after-pubrel").with_clean_start(false), + ); + connect_to(&mut client, addr).await; + let Ok(StagedPublish::Ready(ready)) = client + .stage_publish("t/r".into(), b"r".to_vec(), qos(QoS::ExactlyOnce)) + .await + else { + panic!("publish must stage while connected"); + }; + let packet_id = ready.packet_id().unwrap(); + let claim = claim_quota(&client, &ready).await; + send_with_claim(&client, ready, claim).await; + let after = frames_rx.await.unwrap(); + let disconnect = after.into_iter().find(|f| f[0] == 0xE0); + assert_eq!( + disconnect.and_then(|f| f.get(1).copied()), + Some(0x82), + "an error PUBREC for an identifier already released with PUBREL is a protocol error" + ); + assert_eq!( + client.session.read().await.outbound_stage(packet_id).await, + Some(crate::session::state::OutboundStage::AwaitingPubComp), + "the PUBREL state must be left untouched" + ); + } + + #[tokio::test] + async fn acknowledgement_settled_before_loss_is_reported_as_sent() { + let (completion, handle) = Completion::new(); + completion.delivered(Delivery::AtLeastOnce { packet_id: 9 }); + let (alive, link) = watch::channel(false); + drop(alive); + let settled = InFlight { handle, link }.settle().await; + assert!( + matches!( + settled, + Ok(PublishResult::Sent(Delivery::AtLeastOnce { packet_id: 9 })) + ), + "an acknowledgement that settled before the connection ended must win: {settled:?}" + ); + } + + #[tokio::test] + async fn packet_id_allocation_does_not_scan_the_offline_queue() { + let client = DirectClientInner::new(ConnectOptions::new("offline").with_clean_start(false)); + let mut held: Vec = (1..u16::MAX) + .filter_map(|packet_id| IdReservation::claim(&client.outbound_ids, packet_id)) + .collect(); + assert_eq!(held.len(), usize::from(u16::MAX - 1)); + + let start = std::time::Instant::now(); + let last = client.allocate_packet_id().await; + let exhausted = client.allocate_packet_id().await; + let elapsed = start.elapsed(); + + assert_eq!( + last.as_ref().ok().map(IdReservation::packet_id), + Some(u16::MAX) + ); + assert!(matches!(exhausted, Err(MqttError::PacketIdExhausted))); + assert!( + elapsed < Duration::from_secs(1), + "allocating against a full set of reserved ids took {elapsed:?}" + ); + + held.remove(0); + assert_eq!( + client + .allocate_packet_id() + .await + .ok() + .map(|r| r.packet_id()), + Some(1) + ); + } } diff --git a/crates/mqtt5/src/client/direct/outbound.rs b/crates/mqtt5/src/client/direct/outbound.rs index c47313a0..97840603 100644 --- a/crates/mqtt5/src/client/direct/outbound.rs +++ b/crates/mqtt5/src/client/direct/outbound.rs @@ -45,7 +45,11 @@ impl ServerCapabilities { } pub(crate) fn check_subscribe(self, packet: &SubscribePacket) -> Result<()> { - let identifiers = packet.properties.subscription_identifiers(); + let identifiers = if packet.protocol_version == 5 { + packet.properties.subscription_identifiers() + } else { + Vec::new() + }; if let Some(invalid) = identifiers .iter() .find(|id| !(1..=MAX_SUBSCRIPTION_IDENTIFIER).contains(*id)) @@ -83,7 +87,14 @@ pub(crate) fn check_unsubscribe(packet: &UnsubscribePacket) -> Result<()> { .try_for_each(validate_subscription_filter) } -pub(crate) fn check_publish(topic: &str, options: &PublishOptions) -> Result<()> { +pub(crate) fn check_publish( + topic: &str, + options: &PublishOptions, + protocol_version: u8, +) -> Result<()> { + if protocol_version != 5 { + return validate_topic_name(topic); + } let properties = &options.properties; match (topic.is_empty(), properties.topic_alias) { (true, None) => { @@ -200,15 +211,16 @@ mod tests { #[test] fn publish_validation() { let plain = PublishOptions::default(); - assert!(check_publish("a/b", &plain).is_ok()); - assert!(check_publish("a/+", &plain).is_err()); - assert!(check_publish("", &plain).is_err()); + assert!(check_publish("a/b", &plain, 5).is_ok()); + assert!(check_publish("a/+", &plain, 5).is_err()); + assert!(check_publish("", &plain, 5).is_err()); assert!(check_publish( "", &publish_with(PublishProperties { topic_alias: Some(1), ..Default::default() - }) + }), + 5 ) .is_ok()); assert!(check_publish( @@ -216,7 +228,8 @@ mod tests { &publish_with(PublishProperties { topic_alias: Some(0), ..Default::default() - }) + }), + 5 ) .is_err()); assert!(check_publish( @@ -224,7 +237,8 @@ mod tests { &publish_with(PublishProperties { response_topic: Some("r/#".to_string()), ..Default::default() - }) + }), + 5 ) .is_err()); assert!(check_publish( @@ -232,11 +246,35 @@ mod tests { &publish_with(PublishProperties { subscription_identifiers: vec![1], ..Default::default() - }) + }), + 5 ) .is_err()); } + #[test] + fn v311_publish_skips_v5_property_checks_but_validates_topic() { + let v5_only = publish_with(PublishProperties { + topic_alias: Some(0), + response_topic: Some("r/#".to_string()), + subscription_identifiers: vec![0], + ..Default::default() + }); + assert!(check_publish("a/b", &v5_only, 4).is_ok()); + assert!(check_publish("a/+", &v5_only, 4).is_err()); + assert!(check_publish("", &v5_only, 4).is_err()); + } + + #[test] + fn v311_subscribe_skips_subscription_identifier_checks() { + let mut packet = subscribe("a", Some(0)); + packet.protocol_version = 4; + let no_ids = capabilities(PropertyId::SubscriptionIdentifierAvailable); + assert!(no_ids.check_subscribe(&packet).is_ok()); + packet.filters[0].filter = "a/#/b".to_string(); + assert!(no_ids.check_subscribe(&packet).is_err()); + } + #[test] fn topic_alias_bounds_and_mapping() { let mut aliases = TopicAliasManager::new(2); diff --git a/crates/mqtt5/src/client/direct/reader.rs b/crates/mqtt5/src/client/direct/reader.rs index f9b448d0..97a3c41f 100644 --- a/crates/mqtt5/src/client/direct/reader.rs +++ b/crates/mqtt5/src/client/direct/reader.rs @@ -12,6 +12,7 @@ use crate::packet::unsuback::UnsubAckPacket; use crate::packet::Packet; use crate::protocol::v5::properties::{Properties, PropertyId}; use crate::protocol::v5::reason_codes::ReasonCode; +use crate::session::state::OutboundStage; use crate::session::{SessionState, TopicAliasManager}; use crate::transport::PacketWriter; use parking_lot::Mutex; @@ -30,6 +31,8 @@ use crate::transport::flow::{ FlowFlags, FlowHeader, FlowId, FLOW_TYPE_CLIENT_DATA, FLOW_TYPE_CONTROL, FLOW_TYPE_SERVER_DATA, }; #[cfg(feature = "transport-quic")] +use crate::transport::packet_io::read_packet_from_stream; +#[cfg(feature = "transport-quic")] use bytes::{Buf, Bytes, BytesMut}; #[cfg(feature = "transport-quic")] use quinn::Connection; @@ -44,12 +47,13 @@ pub(super) struct PacketReaderContext { pub(super) callback_manager: Arc, pub(super) suback_channels: Arc>>>, pub(super) unsuback_channels: Arc>>>, - pub(super) puback_channels: Arc>>>, - pub(super) pubcomp_channels: Arc>>>, + pub(super) publish_outcomes: super::tracking::SharedOutcomes, pub(super) writer: Arc>, pub(super) lifecycle: ConnectionLifecycle, #[cfg(feature = "transport-quic")] pub(super) protocol_version: u8, + #[cfg(feature = "transport-quic")] + pub(super) maximum_packet_size: usize, pub(super) auth_handler: Option>, pub(super) auth_method: Option, pub(super) keepalive_state: Arc>, @@ -93,8 +97,6 @@ impl PacketReaderContext { if !self.lifecycle.owns_current_connection() { return; } - self.puback_channels.lock().drain(); - self.pubcomp_channels.lock().drain(); self.suback_channels.lock().drain(); self.unsuback_channels.lock().drain(); } @@ -174,8 +176,10 @@ fn check_problem_information(packet: &Packet, requested: bool) -> Result<()> { pub(super) async fn close_connection( writer: &Arc>, + lifecycle: &ConnectionLifecycle, error: &MqttError, ) { + lifecycle.begin_close(); let disconnect = disconnect_code_for(error) .map(|reason_code| Packet::Disconnect(DisconnectPacket::new(reason_code))); let closing = async { writer.lock().await.close(disconnect).await }; @@ -186,29 +190,46 @@ pub(super) async fn close_connection( } } -async fn process_packet(packet: Packet, ctx: &PacketReaderContext) -> Result<()> { - tracing::trace!("Received packet: {:?}", packet); - check_problem_information(&packet, ctx.request_problem_information)?; - match &packet { - Packet::SubAck(suback) => { - if let Some(tx) = ctx.suback_channels.lock().remove(&suback.packet_id) { - let _ = tx.send(suback.clone()); - return Ok(()); - } - } - Packet::UnsubAck(unsuback) => { - if let Some(tx) = ctx.unsuback_channels.lock().remove(&unsuback.packet_id) { - let _ = tx.send(unsuback.clone()); - return Ok(()); - } +async fn check_acknowledgement_matches(packet: &Packet, ctx: &PacketReaderContext) -> Result<()> { + let (packet_id, name, accepted): (u16, &str, &[OutboundStage]) = match packet { + Packet::PubAck(ack) => (ack.packet_id, "PUBACK", &[OutboundStage::AwaitingPubAck]), + Packet::PubRec(ack) if ack.reason_code.is_error() => { + (ack.packet_id, "PUBREC", &[OutboundStage::AwaitingPubRec]) } + Packet::PubRec(ack) => ( + ack.packet_id, + "PUBREC", + &[ + OutboundStage::AwaitingPubRec, + OutboundStage::AwaitingPubComp, + ], + ), + Packet::PubComp(ack) => (ack.packet_id, "PUBCOMP", &[OutboundStage::AwaitingPubComp]), + _ => return Ok(()), + }; + match ctx.session.read().await.outbound_stage(packet_id).await { + Some(stage) if !accepted.contains(&stage) => Err(MqttError::ProtocolError(format!( + "{name} for packet identifier {packet_id} which is waiting for {stage:?}" + ))), + _ => Ok(()), + } +} + +enum Routed { + Settled, + Unhandled, +} + +async fn settle_acknowledgement(packet: &Packet, ctx: &PacketReaderContext) -> Result { + check_acknowledgement_matches(packet, ctx).await?; + match packet { Packet::PubAck(puback) => { + ctx.publish_outcomes + .lock() + .settle_puback(puback.packet_id, puback.reason_code); super::DirectClientInner::release_outbound_quota(&ctx.session, Some(puback.packet_id)) .await; - if let Some(tx) = ctx.puback_channels.lock().remove(&puback.packet_id) { - let _ = tx.send(puback.reason_code); - return Ok(()); - } + Ok(Routed::Settled) } Packet::PubRec(pubrec) if pubrec.reason_code.is_error() => { tracing::debug!( @@ -216,25 +237,50 @@ async fn process_packet(packet: Packet, ctx: &PacketReaderContext) -> Result<()> reason_code = ?pubrec.reason_code, "QoS 2 PUBREC rejected" ); - if let Some(tx) = ctx.pubcomp_channels.lock().remove(&pubrec.packet_id) { - let _ = tx.send(pubrec.reason_code); - } ctx.session .write() .await .remove_unacked_publish(pubrec.packet_id) .await; + ctx.publish_outcomes.lock().refused( + pubrec.packet_id, + crate::QoS::ExactlyOnce, + pubrec.reason_code, + ); super::DirectClientInner::release_outbound_quota(&ctx.session, Some(pubrec.packet_id)) .await; - return Ok(()); + Ok(Routed::Settled) } Packet::PubComp(pubcomp) => { + ctx.publish_outcomes + .lock() + .acknowledged(pubcomp.packet_id, crate::QoS::ExactlyOnce); super::DirectClientInner::release_outbound_quota(&ctx.session, Some(pubcomp.packet_id)) .await; - if let Some(tx) = ctx.pubcomp_channels.lock().remove(&pubcomp.packet_id) { - let _ = tx.send(pubcomp.reason_code); + Ok(Routed::Settled) + } + _ => Ok(Routed::Unhandled), + } +} + +async fn process_packet(packet: Packet, ctx: &PacketReaderContext) -> Result<()> { + tracing::trace!("Received packet: {:?}", packet); + check_problem_information(&packet, ctx.request_problem_information)?; + if let Routed::Settled = settle_acknowledgement(&packet, ctx).await? { + return Ok(()); + } + match &packet { + Packet::SubAck(suback) => { + if let Some(tx) = ctx.suback_channels.lock().remove(&suback.packet_id) { + let _ = tx.send(suback.clone()); + return Ok(()); + } + } + Packet::UnsubAck(unsuback) => { + if let Some(tx) = ctx.unsuback_channels.lock().remove(&unsuback.packet_id) { + let _ = tx.send(unsuback.clone()); + return Ok(()); } - return Ok(()); } Packet::Auth(auth) => return handle_auth_packet(auth.clone(), ctx).await, _ => {} @@ -261,7 +307,7 @@ pub(super) async fn packet_reader_task_with_responses( }; tracing::error!("Packet reader stopping: {failure}"); - close_connection(&ctx.writer, &failure).await; + close_connection(&ctx.writer, &ctx.lifecycle, &failure).await; drop(reader); ctx.lifecycle.end(disconnect_reason_for(&failure)).await; ctx.clear_pending_if_current(); @@ -487,90 +533,31 @@ async fn try_read_server_flow_header(recv: &mut quinn::RecvStream) -> Result Result { - use crate::packet::FixedHeader; - - while buffer.len() < 2 { - let mut tmp = [0u8; 64]; - let n = recv - .read(&mut tmp) - .await - .map_err(|e| MqttError::ConnectionError(format!("QUIC read error: {e}")))? - .ok_or(MqttError::ClientClosed)?; - if n == 0 { - return Err(MqttError::ClientClosed); - } - buffer.extend_from_slice(&tmp[..n]); - } - - let mut remaining_length = 0u32; - let mut multiplier = 1u32; - let mut remaining_length_bytes = 1usize; - - for i in 1..5 { - if i >= buffer.len() { - let mut tmp = [0u8; 64]; - let n = recv - .read(&mut tmp) - .await - .map_err(|e| MqttError::ConnectionError(format!("QUIC read error: {e}")))? - .ok_or(MqttError::ClientClosed)?; - if n == 0 { - return Err(MqttError::ClientClosed); - } - buffer.extend_from_slice(&tmp[..n]); - } - - let byte = buffer[i]; - remaining_length += u32::from(byte & 0x7F) * multiplier; - multiplier *= 128; - remaining_length_bytes = i; - - if (byte & 0x80) == 0 { - break; - } - - if i == 4 { - return Err(MqttError::MalformedPacket( - "Invalid remaining length encoding".to_string(), - )); - } + ctx: &PacketReaderContext, +) -> Option> { + let read = + read_packet_from_stream(recv, ctx.protocol_version, buffer, ctx.maximum_packet_size).await; + match read { + Ok(_) if ctx.lifecycle.is_closing() => None, + read => Some(read), } +} - let header_len = 1 + remaining_length_bytes; - let total_len = header_len + remaining_length as usize; - - while buffer.len() < total_len { - let needed = total_len - buffer.len(); - let chunk_size = needed.min(4096); - let old_len = buffer.len(); - buffer.resize(old_len + chunk_size, 0); - let n = recv - .read(&mut buffer[old_len..old_len + chunk_size]) - .await - .map_err(|e| MqttError::ConnectionError(format!("QUIC read error: {e}")))? - .ok_or(MqttError::ClientClosed)?; - if n == 0 { - return Err(MqttError::ClientClosed); - } - buffer.truncate(old_len + n); +#[cfg(feature = "transport-quic")] +async fn end_data_stream(ctx: &PacketReaderContext, flow_id: Option, error: &MqttError) { + let fails_connection = + matches!(error, MqttError::ServerDisconnect(_)) || disconnect_code_for(error).is_some(); + if !fails_connection { + tracing::debug!(flow_id = ?flow_id, "Server QUIC data stream closed: {error}"); + return; } - - let packet_bytes = buffer.split_to(total_len); - let mut header_buf = packet_bytes.clone().freeze(); - let fixed_header = FixedHeader::decode(&mut header_buf)?; - - let mut payload_buf = BytesMut::from(&packet_bytes[header_len..]); - Packet::decode_from_body_with_version( - fixed_header.packet_type, - &fixed_header, - &mut payload_buf, - protocol_version, - ) + tracing::error!(flow_id = ?flow_id, "Server QUIC data stream failed the connection: {error}"); + close_connection(&ctx.writer, &ctx.lifecycle, error).await; + ctx.lifecycle.end(disconnect_reason_for(error)).await; + ctx.clear_pending_if_current(); } #[cfg(feature = "transport-quic")] @@ -604,30 +591,31 @@ async fn quic_stream_reader_task( let stream_writer = Arc::new(tokio::sync::Mutex::new(UnifiedWriter::Quic(send))); - loop { - match read_packet_with_buffer(&mut recv, &mut buffer, ctx.protocol_version).await { + while let Some(read) = read_data_stream_packet(&mut recv, &mut buffer, &ctx).await { + let outcome = match read { Ok(packet) => { tracing::trace!(flow_id = ?flow_id, "Received packet on server-initiated QUIC stream: {:?}", packet); - let ack_delivery = ctx.ack_delivery(); - let handlers = ctx.incoming_handlers(ack_delivery.as_ref()); - if let Err(e) = - handle_incoming_packet_with_writer(packet, &stream_writer, flow_id, &handlers) + match settle_acknowledgement(&packet, &ctx).await { + Ok(Routed::Settled) => Ok(()), + Ok(Routed::Unhandled) => { + let ack_delivery = ctx.ack_delivery(); + let handlers = ctx.incoming_handlers(ack_delivery.as_ref()); + handle_incoming_packet_with_writer( + packet, + &stream_writer, + flow_id, + &handlers, + ) .await - { - tracing::error!(flow_id = ?flow_id, "Error handling packet from server stream: {e}"); - if let MqttError::ServerDisconnect(reason_code) = e { - close_connection(&ctx.writer, &e).await; - ctx.lifecycle - .end(DisconnectReason::ServerDisconnect(reason_code)) - .await; } - break; + Err(e) => Err(e), } } - Err(e) => { - tracing::debug!(flow_id = ?flow_id, "Server-initiated QUIC stream closed or error: {e}"); - break; - } + Err(e) => Err(e), + }; + if let Err(e) = outcome { + end_data_stream(&ctx, flow_id, &e).await; + break; } } } @@ -656,33 +644,18 @@ async fn quic_uni_stream_reader_task(mut recv: quinn::RecvStream, ctx: PacketRea } }; - loop { - match read_packet_with_buffer(&mut recv, &mut buffer, ctx.protocol_version).await { + while let Some(read) = read_data_stream_packet(&mut recv, &mut buffer, &ctx).await { + let outcome = match read { Ok(packet) => { tracing::trace!(flow_id = ?flow_id, "Received packet on unidirectional server stream"); - if let Err(e) = handle_incoming_packet_no_writer( - packet, - &ctx.callback_manager, - flow_id, - &ctx.keepalive_state, - ctx.codec_registry.as_ref(), - ) - .await - { - tracing::error!(flow_id = ?flow_id, "Error handling packet from uni stream: {e}"); - if let MqttError::ServerDisconnect(reason_code) = e { - close_connection(&ctx.writer, &e).await; - ctx.lifecycle - .end(DisconnectReason::ServerDisconnect(reason_code)) - .await; - } - break; - } - } - Err(e) => { - tracing::debug!(flow_id = ?flow_id, "Unidirectional server stream closed: {e}"); - break; + handle_incoming_packet_no_writer(packet, flow_id, &ctx.incoming_handlers(None)) + .await } + Err(e) => Err(e), + }; + if let Err(e) = outcome { + end_data_stream(&ctx, flow_id, &e).await; + break; } } } diff --git a/crates/mqtt5/src/client/direct/replay.rs b/crates/mqtt5/src/client/direct/replay.rs index 320050c0..16e5dd96 100644 --- a/crates/mqtt5/src/client/direct/replay.rs +++ b/crates/mqtt5/src/client/direct/replay.rs @@ -1,3 +1,6 @@ +use crate::client::publish_outcome::{ + Delivery, IndeterminateReason, PublishOutcome, PublishRejection, +}; use crate::error::{MqttError, Result}; use crate::packet::publish::PublishPacket; use crate::packet::pubrel::PubRelPacket; @@ -9,9 +12,11 @@ use crate::transport::PacketWriter; use crate::QoS; use parking_lot::Mutex; use std::collections::VecDeque; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::{Arc, Weak}; use tokio::sync::{RwLock, Semaphore}; +use super::tracking::{Completion, IdReservation, SharedIds, SharedOutcomes}; use super::unified::UnifiedWriter; #[derive(Debug, Clone, Copy)] @@ -41,6 +46,120 @@ impl PublishPolicy { } Ok(publish) } + + pub(crate) fn admits_unchanged(self, publish: &PublishPacket) -> Result<()> { + if publish.retain && !self.retain_available { + return Err(MqttError::RetainNotSupported); + } + if self + .maximum_qos + .is_some_and(|maximum| publish.qos as u8 > maximum) + { + return Err(MqttError::QoSNotSupported); + } + Ok(()) + } +} + +#[derive(Debug)] +pub(crate) struct QueuedPublish { + serial: u64, + packet: PublishPacket, + reservation: Option, + completion: Option, +} + +impl QueuedPublish { + pub(crate) fn new( + packet: PublishPacket, + reservation: IdReservation, + completion: Option, + ) -> Self { + Self { + serial: 0, + packet, + reservation: Some(reservation), + completion, + } + } + + fn reject(self, error: &MqttError) { + tracing::warn!( + topic = %self.packet.topic_name, + error = %error, + "Queued message no longer conforms to the connection; not sent" + ); + if let Some(completion) = self.completion { + completion.rejected(PublishRejection::from_error(error)); + } + } +} + +#[derive(Debug, Default)] +pub struct OfflineQueue { + messages: VecDeque, + next_serial: u64, +} + +impl OfflineQueue { + fn stamp(&mut self, mut queued: QueuedPublish) -> QueuedPublish { + self.next_serial = self.next_serial.wrapping_add(1); + queued.serial = self.next_serial; + queued + } + + pub(crate) fn push_back(&mut self, queued: QueuedPublish) { + let queued = self.stamp(queued); + self.messages.push_back(queued); + } + + pub(crate) fn push_front_in_order(&mut self, ordered: Vec) { + for queued in ordered.into_iter().rev() { + let queued = self.stamp(queued); + self.messages.push_front(queued); + } + } + + fn front(&self) -> Option<(u64, PublishPacket)> { + self.messages + .front() + .map(|queued| (queued.serial, queued.packet.clone())) + } + + fn take(&mut self, serial: u64) -> Option { + let position = self + .messages + .iter() + .position(|queued| queued.serial == serial)?; + self.messages.remove(position) + } + + pub(crate) fn is_empty(&self) -> bool { + self.messages.is_empty() + } +} + +#[derive(Clone)] +pub(super) struct ConnectionLink { + pub(super) epoch: u64, + pub(super) current_epoch: Arc, + pub(super) connected: Arc, + pub(super) transfer: Arc>, + pub(super) alive: tokio::sync::watch::Receiver, +} + +impl ConnectionLink { + async fn ended(&self) { + let mut alive = self.alive.clone(); + if alive.wait_for(|alive| !*alive).await.is_err() { + tracing::trace!("Connection liveness signal dropped"); + } + } + + fn is_current(&self) -> bool { + self.connected.load(Ordering::SeqCst) + && self.current_epoch.load(Ordering::SeqCst) == self.epoch + } } pub(super) struct SessionReplay { @@ -48,8 +167,11 @@ pub(super) struct SessionReplay { pub(super) slots: Arc, pub(super) session: Arc>, pub(super) writer: Weak>, - pub(super) queued: Arc>>, + pub(super) queued: Arc>, pub(super) policy: PublishPolicy, + pub(super) outcomes: SharedOutcomes, + pub(super) ids: SharedIds, + pub(super) link: ConnectionLink, } impl SessionReplay { @@ -66,12 +188,23 @@ impl SessionReplay { let packet = match item { OutboundReplay::PubRel(packet_id) => Packet::PubRel(PubRelPacket::new(*packet_id)), OutboundReplay::Publish(publish) => { - if !self.take_slot(flow, publish.packet_id).await { + let resend = without_topic_alias(publish.clone()); + if let Err(e) = self.admits_unchanged(&resend).await { + if !self.abandon_replay(&resend, &e).await { + return false; + } + continue; + } + if !self.take_slot(flow, resend.packet_id).await { return false; } - let mut resend = without_topic_alias(publish.clone()); - resend.dup = true; - Packet::Publish(resend) + if let Some(packet_id) = resend.packet_id { + self.outcomes.lock().mark_resent(packet_id); + } + Packet::Publish(PublishPacket { + dup: true, + ..resend + }) } }; if !self.write(packet).await { @@ -81,51 +214,157 @@ impl SessionReplay { true } + async fn admits_unchanged(&self, publish: &PublishPacket) -> Result<()> { + self.policy.admits_unchanged(publish)?; + self.check_size(publish).await + } + + async fn abandon_replay(&self, publish: &PublishPacket, error: &MqttError) -> bool { + let Some(packet_id) = publish.packet_id else { + return true; + }; + let transfer = self.link.transfer.lock().await; + if !self.link.is_current() { + return false; + } + if publish.qos == QoS::ExactlyOnce { + self.ids.lock().quarantine(packet_id); + } + self.session + .read() + .await + .remove_unacked_publish(packet_id) + .await; + let completion = self.outcomes.lock().take(packet_id); + drop(transfer); + tracing::warn!( + packet_id, + topic = %publish.topic_name, + error = %error, + "Unacknowledged PUBLISH no longer conforms to the resumed connection; not re-sent" + ); + if let Some(completion) = completion { + completion.indeterminate(IndeterminateReason::ReplayNotConforming); + } + true + } + async fn flush_offline_queue(&self, flow: &Arc>) -> bool { loop { - let Some(queued) = self.queued.lock().front().cloned() else { + let Some((serial, queued)) = self.queued.lock().front() else { return true; }; let publish = match self.conform_queued(queued).await { Ok(publish) => publish, Err(e) => { - tracing::warn!("Dropping queued message: {e}"); - self.queued.lock().pop_front(); + let transfer = self.link.transfer.lock().await; + if !self.link.is_current() { + return false; + } + let rejected = self.queued.lock().take(serial); + drop(transfer); + if let Some(rejected) = rejected { + rejected.reject(&e); + } continue; } }; if !self.take_slot(flow, publish.packet_id).await { return false; } - self.queued.lock().pop_front(); - if publish.qos != QoS::AtMostOnce { - if let Err(e) = self - .session - .read() - .await - .store_unacked_publish(publish.clone()) - .await - { - tracing::warn!("Dropping queued message: {e}"); - continue; + match self.transfer_to_session(flow, serial, &publish).await { + Transfer::Stopped => return false, + Transfer::Skipped => {} + Transfer::Stored => { + if !matches!(self.write_publish(publish).await, Written::Complete) { + return false; + } + } + Transfer::Downgraded(queued, transfer) => { + if !self.write_downgraded(publish, queued, transfer).await { + return false; + } } } - if !self.write(Packet::Publish(publish)).await { + } + } + + async fn write_downgraded( + &self, + publish: PublishPacket, + mut queued: QueuedPublish, + transfer: tokio::sync::MutexGuard<'_, ()>, + ) -> bool { + let settled = match self.write_publish(publish).await { + Written::Complete => PublishOutcome::Delivered(Delivery::Unconfirmed), + Written::Failed => PublishOutcome::Indeterminate(IndeterminateReason::ConnectionLost), + Written::NotAttempted => { + tracing::debug!( + "Connection ended before a downgraded message was written; it stays queued" + ); + self.queued.lock().push_front_in_order(vec![queued]); + drop(transfer); return false; } + }; + drop(transfer); + let complete = matches!(settled, PublishOutcome::Delivered(Delivery::Unconfirmed)); + if let Some(completion) = queued.completion.take() { + completion.settle(settled); + } + complete + } + + async fn transfer_to_session( + &self, + flow: &Arc>, + serial: u64, + publish: &PublishPacket, + ) -> Transfer<'_> { + let transfer = self.link.transfer.lock().await; + if !self.link.is_current() { + return Transfer::Stopped; + } + let Some(mut queued) = self.queued.lock().take(serial) else { + drop(transfer); + release_claim(flow, publish.packet_id).await; + return Transfer::Skipped; + }; + let Some(packet_id) = publish.packet_id else { + return Transfer::Downgraded(queued, transfer); + }; + let stored = self + .session + .read() + .await + .store_unacked_publish(publish.clone()) + .await; + if let Err(e) = stored { + drop(transfer); + release_claim(flow, Some(packet_id)).await; + queued.reject(&e); + return Transfer::Skipped; + } + if let Some(completion) = queued.completion.take() { + self.outcomes + .lock() + .track(packet_id, publish.qos, completion); } + queued.reservation.take(); + drop(transfer); + Transfer::Stored } async fn conform_queued(&self, queued: PublishPacket) -> Result { let publish = without_topic_alias(self.policy.conform(queued)?); + self.check_size(&publish).await?; + Ok(publish) + } + + async fn check_size(&self, publish: &PublishPacket) -> Result<()> { let mut buf = bytes::BytesMut::new(); publish.encode(&mut buf)?; - self.session - .read() - .await - .check_packet_size(buf.len()) - .await?; - Ok(publish) + self.session.read().await.check_packet_size(buf.len()).await } async fn take_slot( @@ -143,20 +382,68 @@ impl SessionReplay { .await .claim_send_quota(&self.slots, packet_id) .await + .is_some() } Err(_) => false, } } async fn write(&self, packet: Packet) -> bool { + matches!(self.write_packet(packet).await, Written::Complete) + } + + async fn write_publish(&self, publish: PublishPacket) -> Written { + self.write_packet(Packet::Publish(publish)).await + } + + async fn write_packet(&self, packet: Packet) -> Written { let Some(writer) = self.writer.upgrade() else { - return false; + return Written::NotAttempted; }; - let written = writer.lock().await.write_packet(packet).await; - if let Err(e) = &written { - tracing::debug!("Session replay stopped: {e}"); + let mut writer = tokio::select! { + biased; + () = self.link.ended() => return Written::NotAttempted, + writer = writer.lock() => writer, + }; + if !self.link.is_current() { + tracing::debug!("Session replay stopped: connection replaced"); + return Written::NotAttempted; + } + tokio::select! { + biased; + () = self.link.ended() => { + tracing::debug!("Session replay stopped: connection ended during a write"); + Written::Failed + } + written = writer.write_packet(packet) => match written { + Ok(()) => Written::Complete, + Err(e) => { + tracing::debug!("Session replay stopped: {e}"); + Written::Failed + } + }, + } + } +} + +enum Written { + Complete, + NotAttempted, + Failed, +} + +enum Transfer<'a> { + Stopped, + Skipped, + Stored, + Downgraded(QueuedPublish, tokio::sync::MutexGuard<'a, ()>), +} + +async fn release_claim(flow: &Arc>, packet_id: Option) { + if let Some(packet_id) = packet_id { + if let Err(e) = flow.read().await.acknowledge(packet_id).await { + tracing::trace!(packet_id, "No send quota held: {e}"); } - written.is_ok() } } diff --git a/crates/mqtt5/src/client/direct/tracking.rs b/crates/mqtt5/src/client/direct/tracking.rs new file mode 100644 index 00000000..a91779d9 --- /dev/null +++ b/crates/mqtt5/src/client/direct/tracking.rs @@ -0,0 +1,316 @@ +use crate::client::publish_outcome::{ + Delivery, IndeterminateReason, PublishHandle, PublishOutcome, PublishRejection, +}; +use crate::protocol::v5::reason_codes::ReasonCode; +use crate::QoS; +use parking_lot::Mutex; +use std::collections::hash_map::Entry; +use std::collections::HashMap; +use std::sync::Arc; +use tokio::sync::watch; + +const PACKET_ID_WORDS: usize = (u16::MAX as usize + 1) / 64; + +#[derive(Debug)] +pub(crate) struct PacketIdSet { + words: [u64; PACKET_ID_WORDS], +} + +impl Default for PacketIdSet { + fn default() -> Self { + Self { + words: [0; PACKET_ID_WORDS], + } + } +} + +impl PacketIdSet { + fn bit(packet_id: u16) -> (usize, u64) { + (usize::from(packet_id / 64), 1 << (packet_id % 64)) + } + + fn insert(&mut self, packet_id: u16) { + let (word, bit) = Self::bit(packet_id); + self.words[word] |= bit; + } + + fn remove(&mut self, packet_id: u16) { + let (word, bit) = Self::bit(packet_id); + self.words[word] &= !bit; + } + + fn contains(&self, packet_id: u16) -> bool { + let (word, bit) = Self::bit(packet_id); + self.words[word] & bit != 0 + } + + fn clear(&mut self) { + self.words = [0; PACKET_ID_WORDS]; + } +} + +#[derive(Debug, Default)] +pub(crate) struct OutboundIds { + reserved: PacketIdSet, + quarantined: PacketIdSet, +} + +impl OutboundIds { + pub(crate) fn holds(&self, packet_id: u16) -> bool { + self.reserved.contains(packet_id) || self.quarantined.contains(packet_id) + } + + pub(crate) fn quarantine(&mut self, packet_id: u16) { + self.quarantined.insert(packet_id); + } + + pub(crate) fn release_quarantine(&mut self) { + self.quarantined.clear(); + } +} + +pub(crate) type SharedIds = Arc>; + +#[derive(Debug)] +pub(crate) struct IdReservation { + packet_id: u16, + ids: SharedIds, +} + +impl IdReservation { + pub(crate) fn claim(ids: &SharedIds, packet_id: u16) -> Option { + let mut held = ids.lock(); + if held.holds(packet_id) { + return None; + } + held.reserved.insert(packet_id); + Some(Self { + packet_id, + ids: Arc::clone(ids), + }) + } + + pub(crate) fn packet_id(&self) -> u16 { + self.packet_id + } +} + +impl Drop for IdReservation { + fn drop(&mut self) { + self.ids.lock().reserved.remove(self.packet_id); + } +} + +#[derive(Debug)] +pub(crate) struct Completion { + outcome: watch::Sender>, + resent: bool, +} + +impl Completion { + pub(crate) fn new() -> (Self, PublishHandle) { + let (outcome, handle) = PublishHandle::pending(); + ( + Self { + outcome, + resent: false, + }, + handle, + ) + } + + pub(crate) fn mark_resent(&mut self) { + self.resent = true; + } + + pub(crate) fn settle(self, outcome: PublishOutcome) { + tracing::debug!(?outcome, "publish settled"); + self.outcome.send_replace(Some(outcome)); + } + + pub(crate) fn delivered(self, delivery: Delivery) { + self.settle(PublishOutcome::Delivered(delivery)); + } + + pub(crate) fn rejected(self, rejection: PublishRejection) { + let outcome = match (self.resent, rejection) { + (false, rejection) => PublishOutcome::Rejected(rejection), + (true, PublishRejection::Refused(reason_code)) => { + PublishOutcome::Indeterminate(IndeterminateReason::ResendRefused(reason_code)) + } + (true, _) => PublishOutcome::Indeterminate(IndeterminateReason::ReplayNotConforming), + }; + self.settle(outcome); + } + + pub(crate) fn indeterminate(self, reason: IndeterminateReason) { + self.settle(PublishOutcome::Indeterminate(reason)); + } +} + +#[derive(Debug)] +struct Tracked { + qos: QoS, + completion: Completion, +} + +#[derive(Debug, Default)] +pub(crate) struct OutcomeTracker { + in_flight: HashMap, +} + +pub(crate) type SharedOutcomes = Arc>; + +impl OutcomeTracker { + pub(crate) fn track(&mut self, packet_id: u16, qos: QoS, completion: Completion) { + match self.in_flight.entry(packet_id) { + Entry::Vacant(vacant) => { + vacant.insert(Tracked { qos, completion }); + } + Entry::Occupied(_) => { + tracing::error!( + packet_id, + "packet identifier already tracks an unsettled publish; keeping the existing one" + ); + completion.indeterminate(IndeterminateReason::Abandoned); + } + } + } + + pub(crate) fn take(&mut self, packet_id: u16) -> Option { + self.in_flight + .remove(&packet_id) + .map(|tracked| tracked.completion) + } + + pub(crate) fn mark_resent(&mut self, packet_id: u16) { + if let Some(tracked) = self.in_flight.get_mut(&packet_id) { + tracked.completion.mark_resent(); + } + } + + fn take_matching(&mut self, packet_id: u16, qos: QoS) -> Option { + if self + .in_flight + .get(&packet_id) + .is_some_and(|tracked| tracked.qos == qos) + { + self.take(packet_id) + } else { + None + } + } + + pub(crate) fn acknowledged(&mut self, packet_id: u16, qos: QoS) { + if let Some(completion) = self.take_matching(packet_id, qos) { + completion.delivered(Delivery::of(qos, Some(packet_id))); + } + } + + pub(crate) fn refused(&mut self, packet_id: u16, qos: QoS, reason_code: ReasonCode) { + if let Some(completion) = self.take_matching(packet_id, qos) { + completion.rejected(PublishRejection::Refused(reason_code)); + } + } + + pub(crate) fn settle_puback(&mut self, packet_id: u16, reason_code: ReasonCode) { + if reason_code.is_error() { + self.refused(packet_id, QoS::AtLeastOnce, reason_code); + } else { + self.acknowledged(packet_id, QoS::AtLeastOnce); + } + } + + pub(crate) fn abandon_all(&mut self, reason: IndeterminateReason) { + for (_, tracked) in self.in_flight.drain() { + tracked.completion.indeterminate(reason); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn reservation_is_released_on_drop_and_blocks_duplicates() { + let ids: SharedIds = Arc::default(); + let first = IdReservation::claim(&ids, 7); + assert!(first.is_some()); + assert!(IdReservation::claim(&ids, 7).is_none()); + drop(first); + assert!(IdReservation::claim(&ids, 7).is_some()); + } + + #[test] + fn quarantined_id_cannot_be_reserved_until_released() { + let ids: SharedIds = Arc::default(); + ids.lock().quarantine(9); + assert!(IdReservation::claim(&ids, 9).is_none()); + ids.lock().release_quarantine(); + assert!(IdReservation::claim(&ids, 9).is_some()); + } + + #[test] + fn acknowledgement_settles_only_the_matching_qos() { + let mut tracker = OutcomeTracker::default(); + let (completion, handle) = Completion::new(); + tracker.track(3, QoS::ExactlyOnce, completion); + tracker.settle_puback(3, ReasonCode::Success); + assert_eq!(handle.try_outcome(), None); + tracker.acknowledged(3, QoS::ExactlyOnce); + assert_eq!( + handle.try_outcome(), + Some(PublishOutcome::Delivered(Delivery::ExactlyOnce { + packet_id: 3 + })) + ); + } + + #[test] + fn refusal_after_resend_is_indeterminate() { + let mut tracker = OutcomeTracker::default(); + let (completion, handle) = Completion::new(); + tracker.track(4, QoS::AtLeastOnce, completion); + tracker.mark_resent(4); + tracker.settle_puback(4, ReasonCode::NotAuthorized); + assert_eq!( + handle.try_outcome(), + Some(PublishOutcome::Indeterminate( + IndeterminateReason::ResendRefused(ReasonCode::NotAuthorized) + )) + ); + } + + #[test] + fn tracking_an_occupied_id_keeps_the_existing_publish() { + let mut tracker = OutcomeTracker::default(); + let (first, first_handle) = Completion::new(); + let (second, second_handle) = Completion::new(); + tracker.track(5, QoS::AtLeastOnce, first); + tracker.track(5, QoS::AtLeastOnce, second); + assert_eq!( + second_handle.try_outcome(), + Some(PublishOutcome::Indeterminate( + IndeterminateReason::Abandoned + )) + ); + tracker.settle_puback(5, ReasonCode::Success); + assert_eq!( + first_handle.try_outcome(), + Some(PublishOutcome::Delivered(Delivery::AtLeastOnce { + packet_id: 5 + })) + ); + } + + #[tokio::test] + async fn dropped_completion_resolves_abandoned() { + let (completion, handle) = Completion::new(); + drop(completion); + assert_eq!( + handle.outcome().await, + PublishOutcome::Indeterminate(IndeterminateReason::Abandoned) + ); + } +} diff --git a/crates/mqtt5/src/client/direct/unified.rs b/crates/mqtt5/src/client/direct/unified.rs index d88fa818..e8b41ccd 100644 --- a/crates/mqtt5/src/client/direct/unified.rs +++ b/crates/mqtt5/src/client/direct/unified.rs @@ -12,9 +12,14 @@ use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf}; #[cfg(feature = "transport-websocket")] use crate::transport::websocket::{WebSocketReadHandle, WebSocketWriteHandle}; #[cfg(feature = "transport-quic")] -use crate::transport::PacketReader; +use quinn::{Connection, RecvStream, SendStream}; #[cfg(feature = "transport-quic")] -use quinn::{RecvStream, SendStream}; +use std::sync::Arc; +#[cfg(feature = "transport-quic")] +use std::time::Duration; + +#[cfg(feature = "transport-quic")] +const QUIC_DRAIN_TIMEOUT: Duration = Duration::from_secs(1); enum UnifiedReaderInner { Tcp(OwnedReadHalf), @@ -106,7 +111,15 @@ impl UnifiedReader { .await } #[cfg(feature = "transport-quic")] - UnifiedReaderInner::Quic(reader) => reader.read_packet(self.protocol_version).await, + UnifiedReaderInner::Quic(reader) => { + read_packet_from_stream( + reader, + self.protocol_version, + &mut self.read_buffer, + self.max_packet_size, + ) + .await + } } } } @@ -118,35 +131,66 @@ pub enum UnifiedWriter { WebSocket(WebSocketWriteHandle), #[cfg(feature = "transport-quic")] Quic(SendStream), + #[cfg(feature = "transport-quic")] + QuicControl(SendStream, Arc), Closed, } impl UnifiedWriter { pub async fn close(&mut self, final_packet: Option) -> Result<()> { let mut current = std::mem::replace(self, Self::Closed); + let failed = final_packet.as_ref().is_some_and(is_error_disconnect); let written = match final_packet { Some(packet) => current.write_packet(packet).await, None => Ok(()), }; - let shutdown = current.shutdown().await; + let shutdown = current.shutdown(failed).await; written.and(shutdown) } - async fn shutdown(&mut self) -> Result<()> { + async fn shutdown(&mut self, failed: bool) -> Result<()> { + tracing::debug!(failed, "Shutting down network connection"); match self { Self::Tcp(writer) => Ok(writer.shutdown().await?), Self::Tls(writer) => Ok(writer.shutdown().await?), #[cfg(feature = "transport-websocket")] Self::WebSocket(writer) => writer.close().await, #[cfg(feature = "transport-quic")] - Self::Quic(writer) => writer - .finish() - .map_err(|e| MqttError::ConnectionError(format!("QUIC stream finish: {e}"))), + Self::Quic(writer) => finish_quic_stream(writer), + #[cfg(feature = "transport-quic")] + Self::QuicControl(writer, connection) => { + let finished = finish_quic_stream(writer); + if finished.is_ok() + && tokio::time::timeout(QUIC_DRAIN_TIMEOUT, writer.stopped()) + .await + .is_err() + { + tracing::debug!("QUIC control stream not acknowledged before close"); + } + let code = if failed { + mqtt5_protocol::QuicConnectionCode::Unspecified + } else { + mqtt5_protocol::QuicConnectionCode::NoError + }; + connection.close(quinn::VarInt::from_u32(code.code()), b"disconnect"); + finished + } Self::Closed => Ok(()), } } } +fn is_error_disconnect(packet: &Packet) -> bool { + matches!(packet, Packet::Disconnect(disconnect) if disconnect.reason_code.is_error()) +} + +#[cfg(feature = "transport-quic")] +fn finish_quic_stream(writer: &mut SendStream) -> Result<()> { + writer + .finish() + .map_err(|e| MqttError::ConnectionError(format!("QUIC stream finish: {e}"))) +} + impl PacketWriter for UnifiedWriter { async fn write_packet(&mut self, packet: Packet) -> Result<()> { match self { @@ -155,7 +199,7 @@ impl PacketWriter for UnifiedWriter { #[cfg(feature = "transport-websocket")] Self::WebSocket(writer) => writer.write_packet(packet).await, #[cfg(feature = "transport-quic")] - Self::Quic(writer) => writer.write_packet(packet).await, + Self::Quic(writer) | Self::QuicControl(writer, _) => writer.write_packet(packet).await, Self::Closed => Err(MqttError::NotConnected), } } diff --git a/crates/mqtt5/src/client/mock.rs b/crates/mqtt5/src/client/mock.rs index d1ea2f5b..d3167f41 100644 --- a/crates/mqtt5/src/client/mock.rs +++ b/crates/mqtt5/src/client/mock.rs @@ -9,11 +9,10 @@ use std::sync::atomic::{AtomicBool, AtomicU16, Ordering}; use std::sync::Arc; use tokio::sync::{Mutex, RwLock}; -use crate::client::MqttClientTrait; +use crate::client::{Delivery, MqttClientTrait, PublishResult}; use crate::error::{MqttError, Result}; use crate::types::{ - ConnectOptions, ConnectResult, Message, MessageProperties, PublishOptions, PublishResult, - SubscribeOptions, + ConnectOptions, ConnectResult, Message, MessageProperties, PublishOptions, SubscribeOptions, }; use crate::QoS; @@ -173,7 +172,6 @@ impl MockMqttClient { pub async fn simulate_message(&self, topic: &str, payload: Vec, qos: QoS) -> Result<()> { let subscriptions = self.state.subscriptions.read().await; - // Find matching subscription for (topic_filter, callback) in subscriptions.iter() { if Self::topic_matches(topic_filter, topic) { let message = Message { @@ -200,9 +198,7 @@ impl MockMqttClient { return true; } - // Simple wildcard support for testing if filter.contains('+') || filter.contains('#') { - // Basic implementation - could be enhanced for full MQTT topic matching if filter == "#" { return true; } @@ -210,7 +206,6 @@ impl MockMqttClient { return topic.starts_with(prefix); } if filter.contains('+') { - // Simple single-level wildcard matching let filter_parts: Vec<&str> = filter.split('/').collect(); let topic_parts: Vec<&str> = topic.split('/').collect(); @@ -268,7 +263,6 @@ impl MqttClientTrait for MockMqttClient { } result } else { - // Default behavior: succeed and set connected self.set_connected(true); Ok(()) } @@ -297,7 +291,6 @@ impl MqttClientTrait for MockMqttClient { } result } else { - // Default behavior: succeed and set connected self.set_connected(true); Ok(ConnectResult { session_present: false, @@ -320,7 +313,6 @@ impl MqttClientTrait for MockMqttClient { } result } else { - // Default behavior: succeed and set disconnected self.set_connected(false); Ok(()) } @@ -346,8 +338,7 @@ impl MqttClientTrait for MockMqttClient { if let Some(response) = &responses.publish_response { response.clone() } else { - // Default behavior: succeed with QoS 0 - Ok(PublishResult::QoS0) + Ok(PublishResult::Sent(Delivery::Unconfirmed)) } } } @@ -388,14 +379,16 @@ impl MqttClientTrait for MockMqttClient { if let Some(response) = &responses.publish_response { response.clone() } else { - // Default behavior based on QoS - match options.qos { - QoS::AtMostOnce => Ok(PublishResult::QoS0), - QoS::AtLeastOnce | QoS::ExactlyOnce => { - let packet_id = self.next_packet_id(); - Ok(PublishResult::QoS1Or2 { packet_id }) - } - } + let delivery = match options.qos { + QoS::AtMostOnce => Delivery::Unconfirmed, + QoS::AtLeastOnce => Delivery::AtLeastOnce { + packet_id: self.next_packet_id(), + }, + QoS::ExactlyOnce => Delivery::ExactlyOnce { + packet_id: self.next_packet_id(), + }, + }; + Ok(PublishResult::Sent(delivery)) } } } @@ -416,7 +409,6 @@ impl MqttClientTrait for MockMqttClient { }) .await; - // Store the callback self.state .subscriptions .write() @@ -427,7 +419,6 @@ impl MqttClientTrait for MockMqttClient { if let Some(response) = &responses.subscribe_response { response.clone() } else { - // Default behavior: succeed with packet ID and QoS 0 let packet_id = self.next_packet_id(); Ok((packet_id, QoS::AtMostOnce)) } @@ -452,7 +443,6 @@ impl MqttClientTrait for MockMqttClient { }) .await; - // Store the callback self.state .subscriptions .write() @@ -463,7 +453,6 @@ impl MqttClientTrait for MockMqttClient { if let Some(response) = &responses.subscribe_response { response.clone() } else { - // Default behavior: succeed with packet ID and requested QoS let packet_id = self.next_packet_id(); Ok((packet_id, options.qos)) } @@ -482,14 +471,12 @@ impl MqttClientTrait for MockMqttClient { }) .await; - // Remove the subscription self.state.subscriptions.write().await.remove(&topic_str); let responses = self.state.responses.read().await; if let Some(response) = &responses.unsubscribe_response { response.clone() } else { - // Default behavior: succeed Ok(()) } } diff --git a/crates/mqtt5/src/client/mod.rs b/crates/mqtt5/src/client/mod.rs index 73714610..58b60b78 100644 --- a/crates/mqtt5/src/client/mod.rs +++ b/crates/mqtt5/src/client/mod.rs @@ -10,9 +10,7 @@ use crate::packet::unsubscribe::UnsubscribePacket; use crate::protocol::v5::properties::Properties; #[cfg(not(target_arch = "wasm32"))] use crate::transport::tls::TlsConfig; -use crate::types::{ - ConnectOptions, ConnectResult, PublishOptions, PublishResult, SubscribeOptions, -}; +use crate::types::{ConnectOptions, ConnectResult, PublishOptions, SubscribeOptions}; use crate::QoS; use std::future::Future; use std::sync::Arc; @@ -28,6 +26,7 @@ mod direct; mod error_recovery; mod inner; pub mod mock; +mod publish_outcome; mod retry; mod state; pub mod r#trait; @@ -40,6 +39,9 @@ pub use self::error_recovery::{ is_recoverable, retry_delay, ErrorRecoveryConfig, RecoverableError, RetryState, }; pub use self::mock::{MockCall, MockMqttClient}; +pub use self::publish_outcome::{ + Delivery, IndeterminateReason, PublishHandle, PublishOutcome, PublishRejection, PublishResult, +}; pub use self::r#trait::MqttClientTrait; pub use builders::ConnectionEventCallback; @@ -65,7 +67,7 @@ pub(crate) async fn fire_connection_event( pub use self::direct::AckToken; use self::direct::AutomaticReconnectLifecycle; #[cfg(not(target_arch = "wasm32"))] -use self::direct::{DirectClientInner, StagedPublish}; +use self::direct::{DirectClientInner, ReadyPublish, StagedPublish, Transmitted}; use crate::session::flow_control::FlowControlManager; /// Thread-safe MQTT v5.0 client @@ -397,9 +399,18 @@ impl MqttClient { /// Publishes a message with custom options /// + /// Returns [`PublishResult::Sent`] once a live publish is written (`QoS` 0) or + /// acknowledged (`QoS` 1/2). Returns [`PublishResult::Queued`] when a `QoS` 1/2 + /// publish is queued while disconnected, or when it was sent but the connection + /// ended, the client disconnected or the acknowledgement did not arrive in time; + /// await its handle for the eventual [`crate::PublishOutcome`]. + /// /// # Errors /// - /// Returns an error if the operation fails + /// Returns an error when the publish is definitely not delivered: invalid topic or + /// options, not connected without offline queueing, a capability violation, a + /// failed write before the message was stored, or an error reason code in PUBACK or + /// PUBREC (`PublishFailed`). #[instrument(skip(self, topic, payload, options), fields(qos = ?options.qos, retain = %options.retain))] pub async fn publish_with_options( &self, @@ -427,18 +438,18 @@ impl MqttClient { .stage_publish(topic_str.clone(), payload_vec, options) .await; let outcome = match staged { - Ok(StagedPublish::Queued(result)) => Ok(result), + Ok(StagedPublish::Queued(handle)) => Ok(PublishResult::Queued(handle)), Ok(StagedPublish::Ready(publish)) => self.send_staged_publish(publish).await, Err(e) => Err(e), }; match outcome { Ok(result) => { match &result { - PublishResult::QoS0 => { - tracing::debug!(client_id = %client_id, topic = %topic_str, "Published QoS0 message"); + PublishResult::Sent(delivery) => { + tracing::debug!(client_id = %client_id, topic = %topic_str, ?delivery, "Published message"); } - PublishResult::QoS1Or2 { packet_id } => { - tracing::debug!(client_id = %client_id, topic = %topic_str, packet_id = %packet_id, "Published QoS1/2 message"); + PublishResult::Queued(_) => { + tracing::debug!(client_id = %client_id, topic = %topic_str, "Queued message while disconnected"); } } Ok(result) @@ -450,21 +461,29 @@ impl MqttClient { } } - async fn send_staged_publish(&self, publish: PublishPacket) -> Result { - let packet_id = publish.packet_id; - if let Some(pid) = packet_id { - let flow = Arc::clone(self.inner.read().await.session.read().await.flow_control()); - FlowControlManager::acquire_shared_send_quota(&flow, pid).await?; - } - let ack = self.inner.read().await.transmit_publish(publish).await?; - if let Some(ack) = ack { - ack.wait().await?; + async fn send_staged_publish(&self, mut ready: Box) -> Result { + loop { + let claim = match ready.packet_id() { + Some(pid) => { + let flow = + Arc::clone(self.inner.read().await.session.read().await.flow_control()); + Some(FlowControlManager::acquire_shared_send_quota(&flow, pid).await?) + } + None => None, + }; + let transmitted = self + .inner + .read() + .await + .transmit_publish(ready, claim) + .await?; + match transmitted { + Transmitted::Sent(delivery) => return Ok(PublishResult::Sent(delivery)), + Transmitted::InFlight(in_flight) => return in_flight.settle().await, + Transmitted::Detached(handle) => return Ok(PublishResult::Queued(handle)), + Transmitted::Restaged(next) => ready = next, + } } - Ok( - packet_id.map_or(PublishResult::QoS0, |packet_id| PublishResult::QoS1Or2 { - packet_id, - }), - ) } /// Subscribes to a topic with a callback @@ -555,6 +574,12 @@ impl MqttClient { /// lost error acknowledgement; see [`AckToken::reject`]). Exactly-once applies only /// to a message whose ack completes on a session that survives in memory. /// + /// Acknowledgements leave in arrival order (`[MQTT-4.6.0-2]`/`[MQTT-4.6.0-3]`), so an + /// unresolved token blocks every later PUBACK/PUBREC on the connection, including the + /// automatic ones for plain [`subscribe`](Self::subscribe) callbacks, and holding it + /// keeps those messages in the broker's in-flight window. Resolve tokens promptly; + /// see [`AckToken`] for details. + /// /// # Errors /// Returns an error if `deferred_ack` was not enabled on the connection, or if the /// subscription fails. diff --git a/crates/mqtt5/src/client/publish_outcome.rs b/crates/mqtt5/src/client/publish_outcome.rs new file mode 100644 index 00000000..acd03a6c --- /dev/null +++ b/crates/mqtt5/src/client/publish_outcome.rs @@ -0,0 +1,215 @@ +use crate::protocol::v5::reason_codes::ReasonCode; +use crate::QoS; +use std::future::{Future, IntoFuture}; +use std::pin::Pin; +use tokio::sync::watch; + +/// How a publish reached the server. +/// +/// `QoS` 0 has no acknowledgement, so a message that went out at `QoS` 0 (requested, or +/// downgraded to the server's Maximum `QoS` 0) is [`Delivery::Unconfirmed`]: it was +/// written to the connection, but the server never confirms receipt. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Delivery { + /// Written at `QoS` 0. Receipt is not confirmed by the server. + Unconfirmed, + /// Sent at `QoS` 1 and acknowledged with a successful PUBACK. + AtLeastOnce { + /// Packet identifier the acknowledged PUBLISH carried. + packet_id: u16, + }, + /// Sent at `QoS` 2 and completed with PUBCOMP. + ExactlyOnce { + /// Packet identifier the completed PUBLISH carried. + packet_id: u16, + }, +} + +impl Delivery { + /// The `QoS` the message was actually sent with. It is lower than the requested + /// `QoS` when the server's Maximum `QoS` forced a downgrade. + #[must_use] + pub fn qos_used(self) -> QoS { + match self { + Self::Unconfirmed => QoS::AtMostOnce, + Self::AtLeastOnce { .. } => QoS::AtLeastOnce, + Self::ExactlyOnce { .. } => QoS::ExactlyOnce, + } + } + + /// Packet identifier of an acknowledged delivery. + #[must_use] + pub fn packet_id(self) -> Option { + match self { + Self::Unconfirmed => None, + Self::AtLeastOnce { packet_id } | Self::ExactlyOnce { packet_id } => Some(packet_id), + } + } + + pub(crate) fn of(qos: QoS, packet_id: Option) -> Self { + match (qos, packet_id) { + (QoS::AtLeastOnce, Some(packet_id)) => Self::AtLeastOnce { packet_id }, + (QoS::ExactlyOnce, Some(packet_id)) => Self::ExactlyOnce { packet_id }, + _ => Self::Unconfirmed, + } + } +} + +/// Why a publish was definitely not delivered. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[non_exhaustive] +pub enum PublishRejection { + /// The message is retained but the server reported Retain Available 0. + RetainNotSupported, + /// The encoded PUBLISH exceeds the server's Maximum Packet Size. + PacketTooLarge, + /// The server does not support the requested `QoS`. + QoSNotSupported, + /// The server refused the message with this reason code in PUBACK or PUBREC. + Refused(ReasonCode), + /// The client could not encode or record the message. + Unsendable, +} + +impl PublishRejection { + pub(crate) fn from_error(error: &crate::error::MqttError) -> Self { + match error { + crate::error::MqttError::RetainNotSupported => Self::RetainNotSupported, + crate::error::MqttError::PacketTooLarge { .. } => Self::PacketTooLarge, + crate::error::MqttError::QoSNotSupported => Self::QoSNotSupported, + crate::error::MqttError::PublishFailed(reason_code) => Self::Refused(*reason_code), + _ => Self::Unsendable, + } + } +} + +/// Why the delivery of a publish cannot be determined: it may or may not have reached +/// the server. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[non_exhaustive] +pub enum IndeterminateReason { + /// The client asked to resume its session (Clean Start 0) but the server did not + /// (Session Present 0) while a `QoS` 2 exchange was unacknowledged. Unacknowledged + /// `QoS` 1 messages are re-sent instead. + SessionLost, + /// The client connected with Clean Start 1 and, as MQTT-3.1.2-4 requires, + /// discarded its session state while the exchange was unacknowledged. + SessionDiscarded, + /// A message that had already been sent no longer conforms to the capabilities of + /// the new connection (Retain Available, Maximum Packet Size, Maximum `QoS`), so it + /// was not re-sent. + ReplayNotConforming, + /// A message that had already been sent was refused when it was re-sent. + ResendRefused(ReasonCode), + /// The connection failed while a message downgraded to `QoS` 0 was being written. + ConnectionLost, + /// The client was dropped before the publish settled. + Abandoned, +} + +/// Final outcome of a publish. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PublishOutcome { + /// The message reached the server; see [`Delivery`] for the `QoS` used. + Delivered(Delivery), + /// The message was definitely not delivered. + Rejected(PublishRejection), + /// The message may have been delivered. + Indeterminate(IndeterminateReason), +} + +/// Pending outcome of a publish that has not settled yet: one accepted into the offline +/// queue, or one sent on a connection that ended (or did not acknowledge it in time) +/// before it was acknowledged. +/// +/// Await the handle (it implements [`IntoFuture`]) or call [`PublishHandle::outcome`] +/// to wait for the single [`PublishOutcome`] of that publish. The outcome belongs to the +/// publish itself, never to a packet identifier, so reuse of identifiers cannot settle +/// the wrong publish. Clones observe the same outcome. +/// +/// The handle stays pending while the client is disconnected. It settles once the client +/// reconnects and the message is acknowledged, rejected or found to be lost with the +/// session; if the client is dropped first it resolves +/// [`IndeterminateReason::Abandoned`]. +#[derive(Debug, Clone)] +pub struct PublishHandle { + outcome: watch::Receiver>, +} + +impl PublishHandle { + /// A handle that is already settled with `outcome`. + #[must_use] + pub fn resolved(outcome: PublishOutcome) -> Self { + let (_, outcome) = watch::channel(Some(outcome)); + Self { outcome } + } + + pub(crate) fn pending() -> (watch::Sender>, Self) { + let (sender, outcome) = watch::channel(None); + (sender, Self { outcome }) + } + + /// The outcome if the publish has already settled. + #[must_use] + pub fn try_outcome(&self) -> Option { + *self.outcome.borrow() + } + + /// Waits until the publish settles. + pub async fn outcome(mut self) -> PublishOutcome { + match self.outcome.wait_for(Option::is_some).await { + Ok(settled) => settled.unwrap_or(PublishOutcome::Indeterminate( + IndeterminateReason::Abandoned, + )), + Err(_) => PublishOutcome::Indeterminate(IndeterminateReason::Abandoned), + } + } +} + +impl IntoFuture for PublishHandle { + type Output = PublishOutcome; + type IntoFuture = Pin + Send>>; + + fn into_future(self) -> Self::IntoFuture { + Box::pin(self.outcome()) + } +} + +/// Result of a publish call. +/// +/// A publish on a live connection normally settles before the call returns: `QoS` 0 +/// once written, `QoS` 1/2 once acknowledged ([`PublishResult::Sent`]). +/// +/// [`PublishResult::Queued`] means the publish has not settled yet and its handle +/// yields the eventual outcome. That happens when a `QoS` 1/2 publish is made while +/// disconnected with offline queueing enabled, and also when a `QoS` 1/2 publish was +/// sent but the connection ended, the client disconnected, or no acknowledgement +/// arrived in time: the message stays in flight with the session and is re-sent on the +/// next connection. The handle stays pending until the client reconnects (then it +/// settles) or is dropped (then it resolves [`IndeterminateReason::Abandoned`]). +#[derive(Debug, Clone)] +pub enum PublishResult { + /// Sent on the live connection. + Sent(Delivery), + /// Queued offline, or in flight across a connection loss; the handle settles later. + Queued(PublishHandle), +} + +impl PublishResult { + /// Packet identifier of a publish acknowledged on the live connection. + #[must_use] + pub fn packet_id(&self) -> Option { + match self { + Self::Sent(delivery) => delivery.packet_id(), + Self::Queued(_) => None, + } + } + + /// Waits for the final outcome of either kind of publish. + pub async fn outcome(self) -> PublishOutcome { + match self { + Self::Sent(delivery) => PublishOutcome::Delivered(delivery), + Self::Queued(handle) => handle.outcome().await, + } + } +} diff --git a/crates/mqtt5/src/client/trait.rs b/crates/mqtt5/src/client/trait.rs index 4b7ecb9a..3d13a82b 100644 --- a/crates/mqtt5/src/client/trait.rs +++ b/crates/mqtt5/src/client/trait.rs @@ -3,10 +3,9 @@ //! This module provides a trait interface for the MQTT client to enable //! mocking and testing without a real broker connection. +use crate::client::PublishResult; use crate::error::Result; -use crate::types::{ - ConnectOptions, ConnectResult, Message, PublishOptions, PublishResult, SubscribeOptions, -}; +use crate::types::{ConnectOptions, ConnectResult, Message, PublishOptions, SubscribeOptions}; use crate::QoS; use std::future::Future; diff --git a/crates/mqtt5/src/lib.rs b/crates/mqtt5/src/lib.rs index c7f7a90b..5c96a36b 100644 --- a/crates/mqtt5/src/lib.rs +++ b/crates/mqtt5/src/lib.rs @@ -246,8 +246,9 @@ pub mod types; #[cfg(not(target_arch = "wasm32"))] pub use client::{ - AckToken, AuthHandler, AuthResponse, ConnectionEvent, DisconnectReason, JwtAuthHandler, - MockCall, MockMqttClient, MqttClient, MqttClientTrait, PlainAuthHandler, + AckToken, AuthHandler, AuthResponse, ConnectionEvent, Delivery, DisconnectReason, + IndeterminateReason, JwtAuthHandler, MockCall, MockMqttClient, MqttClient, MqttClientTrait, + PlainAuthHandler, PublishHandle, PublishOutcome, PublishRejection, PublishResult, ScramSha256AuthHandler, }; #[cfg(feature = "codec-deflate")] @@ -261,7 +262,7 @@ pub use mqtt5_protocol::{ validate_client_id, validate_topic_filter, validate_topic_name, ConnectProperties, ConnectResult, FixedHeader, Message, MessageProperties, MqttError, Packet, PacketType, Properties, PropertyId, PropertyValue, PropertyValueType, ProtocolVersion, PublishOptions, - PublishProperties, PublishResult, QoS, RestrictiveValidator, Result, RetainHandling, - StandardValidator, SubscribeOptions, TopicValidator, Transport, WillMessage, WillProperties, + PublishProperties, QoS, RestrictiveValidator, Result, RetainHandling, StandardValidator, + SubscribeOptions, TopicValidator, Transport, WillMessage, WillProperties, }; pub use types::{ConnectOptions, ConnectionStats}; diff --git a/crates/mqtt5/src/session/flow_control.rs b/crates/mqtt5/src/session/flow_control.rs index f97e25b8..eafd87f1 100644 --- a/crates/mqtt5/src/session/flow_control.rs +++ b/crates/mqtt5/src/session/flow_control.rs @@ -29,6 +29,7 @@ pub struct FlowControlManager { inbound_in_flight: Arc>>, quota_debt: Arc, replay_slots: Arc>>>, + quota_generation: u64, } /// A pending publish request waiting for quota @@ -69,6 +70,7 @@ impl FlowControlManager { inbound_in_flight: Arc::new(RwLock::new(HashMap::new())), quota_debt: Arc::new(AtomicUsize::new(0)), replay_slots: Arc::new(Mutex::new(None)), + quota_generation: 0, } } @@ -127,6 +129,7 @@ impl FlowControlManager { drop(in_flight); self.receive_maximum = receive_maximum; + self.quota_generation = self.quota_generation.wrapping_add(1); let capacity = if receive_maximum == 0 { Semaphore::MAX_PERMITS } else { @@ -175,7 +178,16 @@ impl FlowControlManager { drop(in_flight); } - pub async fn claim_send_quota(&self, semaphore: &Arc, packet_id: u16) -> bool { + #[must_use] + pub fn quota_generation(&self) -> u64 { + self.quota_generation + } + + pub async fn claim_send_quota( + &self, + semaphore: &Arc, + packet_id: u16, + ) -> Option { let issued_by_current = Arc::ptr_eq(semaphore, &self.quota_semaphore) || self .replay_slots @@ -188,19 +200,22 @@ impl FlowControlManager { .await .insert(packet_id, Instant::now()); } - issued_by_current + issued_by_current.then_some(self.quota_generation) } /// # Errors /// /// Returns `FlowControlExceeded` when the backpressure timeout elapses and /// `NotConnected` when the quota was closed by a disconnect. - pub async fn acquire_shared_send_quota(flow: &Arc>, packet_id: u16) -> Result<()> { + pub async fn acquire_shared_send_quota( + flow: &Arc>, + packet_id: u16, + ) -> Result { loop { let (semaphore, timeout) = { let manager = flow.read().await; if manager.receive_maximum == 0 { - return Ok(()); + return Ok(manager.quota_generation); } ( Arc::clone(&manager.quota_semaphore), @@ -215,13 +230,13 @@ impl FlowControlManager { }; if let Ok(permit) = acquired { permit.forget(); - if flow + if let Some(generation) = flow .read() .await .claim_send_quota(&semaphore, packet_id) .await { - return Ok(()); + return Ok(generation); } } else if Arc::ptr_eq(&semaphore, &flow.read().await.quota_semaphore) { return Err(MqttError::NotConnected); @@ -552,7 +567,12 @@ mod tests { assert_eq!(slots.available_permits(), 1); slots.acquire().await.unwrap().forget(); - assert!(flow.read().await.claim_send_quota(&slots, 7).await); + assert!(flow + .read() + .await + .claim_send_quota(&slots, 7) + .await + .is_some()); flow.read().await.acknowledge(7).await.unwrap(); assert_eq!(slots.available_permits(), 1); assert_eq!(flow.read().await.available_permits(), 0); @@ -583,7 +603,12 @@ mod tests { let stale = Arc::clone(&flow.read().await.quota_semaphore); flow.write().await.reset_for_connection(2, &[], false).await; assert!(stale.is_closed()); - assert!(!flow.read().await.claim_send_quota(&stale, 1).await); + assert!(flow + .read() + .await + .claim_send_quota(&stale, 1) + .await + .is_none()); } #[tokio::test] diff --git a/crates/mqtt5/src/session/state.rs b/crates/mqtt5/src/session/state.rs index da31cbd8..9d7b5e2f 100644 --- a/crates/mqtt5/src/session/state.rs +++ b/crates/mqtt5/src/session/state.rs @@ -96,6 +96,17 @@ pub struct SessionState { flow_registry: Arc>, } +/// Which acknowledgement an outbound packet identifier is waiting for. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum OutboundStage { + /// A `QoS` 1 PUBLISH waiting for PUBACK. + AwaitingPubAck, + /// A `QoS` 2 PUBLISH waiting for PUBREC. + AwaitingPubRec, + /// A `QoS` 2 PUBREL waiting for PUBCOMP. + AwaitingPubComp, +} + #[derive(Debug, Clone)] pub enum OutboundReplay { Publish(PublishPacket), @@ -355,6 +366,26 @@ impl SessionState { ordered.into_iter().map(|(_, item)| item).collect() } + /// Reports which acknowledgement the outbound packet identifier is waiting for. + pub async fn outbound_stage(&self, packet_id: u16) -> Option { + let publish_qos = self + .unacked_publishes + .read() + .await + .get(&packet_id) + .map(|(_, publish)| publish.qos); + match publish_qos { + Some(crate::QoS::ExactlyOnce) => Some(OutboundStage::AwaitingPubRec), + Some(_) => Some(OutboundStage::AwaitingPubAck), + None => self + .unacked_pubrels + .read() + .await + .contains_key(&packet_id) + .then_some(OutboundStage::AwaitingPubComp), + } + } + pub async fn discard_outbound_state(&self) { self.unacked_publishes.write().await.clear(); self.unacked_pubrels.write().await.clear(); diff --git a/crates/mqtt5/src/types.rs b/crates/mqtt5/src/types.rs index 67210f02..3f3b66ee 100644 --- a/crates/mqtt5/src/types.rs +++ b/crates/mqtt5/src/types.rs @@ -18,6 +18,11 @@ pub struct ConnectOptions { /// Enable deferred acknowledgement: inbound `QoS` > 0 messages are delivered with /// an `AckToken` the application resolves after durable processing. Requires a /// persistent session and a bounded receive maximum; see `validate_deferred_ack`. + /// + /// Acknowledgements are sent in arrival order (`[MQTT-4.6.0-2]`/`[MQTT-4.6.0-3]`), so + /// an unresolved `AckToken` also holds back the automatic PUBACK/PUBREC of every later + /// message, including those for plain `subscribe` callbacks, and can stall the + /// broker's in-flight window. See the `AckToken` head-of-line blocking notes. pub deferred_ack: bool, /// Accept a broker-held session that this client instance has no local state for. /// @@ -63,6 +68,12 @@ impl ConnectOptions { /// Enables deferred acknowledgement. Connection-wide: every `subscribe_with_ack` /// subscription delivers an `AckToken`. Must be paired with a persistent session /// and an explicit bounded receive maximum (see `validate_deferred_ack`). + /// + /// Once enabled, an unresolved `AckToken` blocks all later PUBACK/PUBREC on the + /// connection, including automatic acknowledgements for plain `subscribe` + /// callbacks, because acknowledgements must follow arrival order + /// (`[MQTT-4.6.0-2]`/`[MQTT-4.6.0-3]`). A long-held token can therefore exhaust the + /// Receive Maximum window and stall all inbound `QoS` 1/2 delivery. #[must_use] pub fn with_deferred_ack(mut self, on: bool) -> Self { self.deferred_ack = on; @@ -255,8 +266,8 @@ impl DerefMut for ConnectOptions { pub use mqtt5_protocol::{ ConnectProperties, ConnectResult, KeepaliveConfig, Message, MessageProperties, ProtocolVersion, - PublishOptions, PublishProperties, PublishResult, RetainHandling, SubscribeOptions, - WillMessage, WillProperties, + PublishOptions, PublishProperties, RetainHandling, SubscribeOptions, WillMessage, + WillProperties, }; #[derive(Debug, Clone, Default)] diff --git a/crates/mqtt5/tests/broker_bridge_integration.rs b/crates/mqtt5/tests/broker_bridge_integration.rs index b37d4b0b..570e6e79 100644 --- a/crates/mqtt5/tests/broker_bridge_integration.rs +++ b/crates/mqtt5/tests/broker_bridge_integration.rs @@ -2,10 +2,9 @@ //! Integration tests for broker-to-broker bridging use mqtt5::broker::bridge::{BridgeConfig, BridgeDirection, BridgeManager}; -use mqtt5::broker::router::MessageRouter; +use mqtt5::broker::router::{MessageRouter, SubscriptionRequest}; use mqtt5::packet::publish::PublishPacket; use mqtt5::time::Duration; -use mqtt5::types::ProtocolVersion; use mqtt5::QoS; use std::sync::Arc; @@ -96,18 +95,11 @@ async fn test_bridge_message_routing() { ) .await; router - .subscribe( - "test-client".to_string(), - "test/topic".to_string(), + .subscribe(SubscriptionRequest::new( + "test-client", + "test/topic", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); diff --git a/crates/mqtt5/tests/change_only_delivery.rs b/crates/mqtt5/tests/change_only_delivery.rs index a03320e6..9861a647 100644 --- a/crates/mqtt5/tests/change_only_delivery.rs +++ b/crates/mqtt5/tests/change_only_delivery.rs @@ -1,9 +1,8 @@ #![cfg(feature = "broker")] -use mqtt5::broker::router::MessageRouter; +use mqtt5::broker::router::{MessageRouter, SubscriptionRequest}; use mqtt5::broker::storage::ChangeOnlyState; use mqtt5::packet::publish::PublishPacket; use mqtt5::time::Duration; -use mqtt5::types::ProtocolVersion; use mqtt5::QoS; use std::sync::Arc; use tokio::time::timeout; @@ -26,16 +25,8 @@ async fn test_change_only_filters_duplicate_payloads() { router .subscribe( - "change_only_client".to_string(), - "sensors/temperature".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("change_only_client", "sensors/temperature", QoS::AtMostOnce) + .with_change_only(true), ) .await .unwrap(); @@ -82,16 +73,8 @@ async fn test_change_only_allows_different_payloads() { router .subscribe( - "change_only_client".to_string(), - "sensors/temperature".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("change_only_client", "sensors/temperature", QoS::AtMostOnce) + .with_change_only(true), ) .await .unwrap(); @@ -138,18 +121,11 @@ async fn test_change_only_disabled_allows_duplicates() { .await; router - .subscribe( - "regular_client".to_string(), - "sensors/temperature".to_string(), + .subscribe(SubscriptionRequest::new( + "regular_client", + "sensors/temperature", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -195,16 +171,8 @@ async fn test_change_only_per_topic_tracking() { router .subscribe( - "change_only_client".to_string(), - "sensors/+".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("change_only_client", "sensors/+", QoS::AtMostOnce) + .with_change_only(true), ) .await .unwrap(); @@ -251,16 +219,8 @@ async fn test_change_only_state_persistence() { router .subscribe( - "persistent_client".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("persistent_client", "test/topic", QoS::AtMostOnce) + .with_change_only(true), ) .await .unwrap(); @@ -308,16 +268,8 @@ async fn test_change_only_state_load_on_reconnect() { router .subscribe( - "reconnect_client".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("reconnect_client", "test/topic", QoS::AtMostOnce) + .with_change_only(true), ) .await .unwrap(); @@ -399,16 +351,8 @@ async fn test_change_only_with_qos_levels() { router .subscribe( - "qos_client".to_string(), - "test/topic".to_string(), - QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("qos_client", "test/topic", QoS::AtLeastOnce) + .with_change_only(true), ) .await .unwrap(); @@ -466,32 +410,16 @@ async fn test_change_only_multiple_clients_independent() { router .subscribe( - "client1".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("client1", "test/topic", QoS::AtMostOnce) + .with_change_only(true), ) .await .unwrap(); router .subscribe( - "client2".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - true, - None, + SubscriptionRequest::new("client2", "test/topic", QoS::AtMostOnce) + .with_change_only(true), ) .await .unwrap(); diff --git a/crates/mqtt5/tests/client_publish.rs b/crates/mqtt5/tests/client_publish.rs index 40b443f9..f9f77cbe 100644 --- a/crates/mqtt5/tests/client_publish.rs +++ b/crates/mqtt5/tests/client_publish.rs @@ -1,4 +1,4 @@ -use mqtt5::{MqttClient, MqttError, PublishOptions, PublishResult, QoS}; +use mqtt5::{Delivery, MqttClient, MqttError, PublishOptions, PublishResult, QoS}; use ulid::Ulid; /// Generate a lexicographically sortable client ID using ULID @@ -18,16 +18,16 @@ async fn test_publish_not_connected() { async fn test_publish_qos0() { let client = MqttClient::new(test_client_id("publish-qos0")); - // Try connecting match client.connect("mqtt://127.0.0.1:1883").await { Ok(()) => { - // Test QoS 0 publish let result = client.publish_qos0("test/topic", b"QoS 0 message").await; assert!(result.is_ok()); - // QoS 0 doesn't return a packet ID let result = client.publish("test/topic", b"Another QoS 0").await; - assert!(matches!(result, Ok(PublishResult::QoS0))); + assert!(matches!( + result, + Ok(PublishResult::Sent(Delivery::Unconfirmed)) + )); client.disconnect().await.unwrap(); } @@ -43,12 +43,13 @@ async fn test_publish_qos1() { match client.connect("mqtt://127.0.0.1:1883").await { Ok(()) => { - // Test QoS 1 publish - should return packet ID let result = client.publish_qos1("test/qos1", b"QoS 1 message").await; assert!(result.is_ok()); match result.unwrap() { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => assert!(packet_id > 0), + other => panic!("expected an acknowledged publish, got {other:?}"), } client.disconnect().await.unwrap(); @@ -65,12 +66,13 @@ async fn test_publish_qos2() { match client.connect("mqtt://127.0.0.1:1883").await { Ok(()) => { - // Test QoS 2 publish - should return packet ID let result = client.publish_qos2("test/qos2", b"QoS 2 message").await; assert!(result.is_ok()); match result.unwrap() { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => assert!(packet_id > 0), + other => panic!("expected an acknowledged publish, got {other:?}"), } client.disconnect().await.unwrap(); @@ -87,7 +89,6 @@ async fn test_publish_retain() { match client.connect("mqtt://127.0.0.1:1883").await { Ok(()) => { - // Test retained message let result = client .publish_retain("test/retained", b"Retained message") .await; @@ -107,7 +108,6 @@ async fn test_publish_with_options() { match client.connect("mqtt://127.0.0.1:1883").await { Ok(()) => { - // Test publish with custom options let mut options = PublishOptions { qos: QoS::AtLeastOnce, retain: true, @@ -125,8 +125,10 @@ async fn test_publish_with_options() { assert!(result.is_ok()); match result.unwrap() { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result for QoS 1 publish"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => assert!(packet_id > 0), + other => panic!("expected an acknowledged publish, got {other:?}"), } client.disconnect().await.unwrap(); @@ -143,7 +145,6 @@ async fn test_publish_empty_payload() { match client.connect("mqtt://127.0.0.1:1883").await { Ok(()) => { - // Empty payload is valid let result = client.publish("test/empty", b"").await; assert!(result.is_ok()); @@ -161,7 +162,6 @@ async fn test_publish_string_conversion() { match client.connect("mqtt://127.0.0.1:1883").await { Ok(()) => { - // Test string conversions let result = client.publish("test/string", "String payload").await; assert!(result.is_ok()); diff --git a/crates/mqtt5/tests/conf_client_a.rs b/crates/mqtt5/tests/conf_client_a.rs index 2a1919e5..dcf16ee0 100644 --- a/crates/mqtt5/tests/conf_client_a.rs +++ b/crates/mqtt5/tests/conf_client_a.rs @@ -1258,7 +1258,7 @@ async fn mqtt_2_1_3_1_reserved_flags_on_inbound_puback_is_malformed() { .unwrap(); let r = timeout(Duration::from_secs(2), h).await.unwrap().unwrap(); assert!( - r.is_err() && !s.client.is_connected().await, + !matches!(r, Ok(mqtt5::PublishResult::Sent(_))) && !s.client.is_connected().await, "MQTT-2.1.3-1 VIOLATION: PUBACK with reserved flag bits 0x2 accepted as a valid acknowledgement ({r:?})" ); } diff --git a/crates/mqtt5/tests/conf_client_b.rs b/crates/mqtt5/tests/conf_client_b.rs index 7dd3691d..2f68ed31 100644 --- a/crates/mqtt5/tests/conf_client_b.rs +++ b/crates/mqtt5/tests/conf_client_b.rs @@ -6,8 +6,8 @@ use mqtt5::packet::{FixedHeader, MqttPacket, Packet, PacketType}; use mqtt5::protocol::v5::reason_codes::ReasonCode; use mqtt5::session::TopicAliasManager; use mqtt5::{ - AckToken, ConnectOptions, Message, MqttClient, MqttError, PublishOptions, PublishProperties, - QoS, SubscribeOptions, + AckToken, ConnectOptions, Message, MqttClient, MqttError, ProtocolVersion, PublishOptions, + PublishProperties, QoS, SubscribeOptions, }; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -867,6 +867,14 @@ async fn mqtt_3_3_2_10_client_accepts_topic_alias_within_its_maximum() { } async fn alias_violation(alias: u16, client_id: &str) -> (usize, Termination) { + alias_violation_on_topic("t/a", alias, client_id).await +} + +async fn alias_violation_on_topic( + topic: &str, + alias: u16, + client_id: &str, +) -> (usize, Termination) { let (client, mut stream, _listener) = connect_client(alias_client_options(client_id), &plain_connack()).await; let mut rx = subscribe(&client, &mut stream, "t/a", qos_options(QoS::AtMostOnce)).await; @@ -874,7 +882,7 @@ async fn alias_violation(alias: u16, client_id: &str) -> (usize, Termination) { .write_all(&raw_publish( 0, false, - "t/a", + topic, None, &topic_alias_prop(alias), b"x", @@ -903,6 +911,24 @@ async fn mqtt_3_3_2_11_receiver_treats_topic_alias_above_its_maximum_as_protocol ); } +#[tokio::test] +async fn mqtt_3_3_2_8_inbound_topic_alias_zero_with_empty_topic_is_topic_alias_invalid() { + let (delivered, end) = alias_violation_on_topic("", 0, "conf-b-alias0-empty").await; + assert!( + delivered == 0 && end == Termination::Disconnect(0x94), + "[MQTT-3.3.2-8 / §3.3.2.3.4] inbound Topic Alias 0 with a zero-length Topic Name must draw DISCONNECT 0x94 (Topic Alias invalid), not 0x82; delivered={delivered} termination={end:?}" + ); +} + +#[tokio::test] +async fn mqtt_3_3_2_11_inbound_topic_alias_above_maximum_with_empty_topic_is_topic_alias_invalid() { + let (delivered, end) = alias_violation_on_topic("", 3, "conf-b-alias3-empty").await; + assert!( + delivered == 0 && end == Termination::Disconnect(0x94), + "[MQTT-3.3.2-11 / §3.3.2.3.4] inbound Topic Alias 3 > client maximum 2 with a zero-length Topic Name must draw DISCONNECT 0x94 (Topic Alias invalid), not 0x82; delivered={delivered} termination={end:?}" + ); +} + #[tokio::test] async fn crosscheck3_inbound_subscription_identifier_zero_is_protocol_error() { let (client, mut stream, _listener) = @@ -1130,8 +1156,8 @@ async fn mqtt_4_9_0_1_send_quota_reinitialized_on_new_connection() { .publish_qos1("t/quota", b"never-acked".to_vec()) .await; assert!( - matches!(timed_out, Err(MqttError::Timeout)), - "unacknowledged publish times out: {timed_out:?}" + matches!(&timed_out, Ok(mqtt5::PublishResult::Queued(handle)) if handle.try_outcome().is_none()), + "unacknowledged publish returns a pending handle after the ack wait: {timed_out:?}" ); let _ = collect_frames(&mut first, Duration::from_millis(50)).await; drop(first); @@ -1179,7 +1205,7 @@ async fn queued_flush(client_id: &str) -> Vec { .publish_qos1("t/queued", vec![i]) .await .expect("publish while disconnected is queued"); - assert!(queued.packet_id().is_some()); + assert!(matches!(queued, mqtt5::PublishResult::Queued(_))); } let (mut second, _) = accept_session( @@ -1507,3 +1533,44 @@ async fn mqtt_3_2_2_5_queued_acks_discarded_when_session_not_present() { "[MQTT-3.2.2-5] acknowledgements from the discarded session were sent on a Session Present=0 connection: {pubacks:?}" ); } + +#[tokio::test] +async fn mqtt_3_1_1_publish_ignores_v5_only_properties_and_sends_none() { + let (listener, url) = bind().await; + let client = MqttClient::with_options( + base_options("conf-b-v311-props").with_protocol_version(ProtocolVersion::V311), + ); + let connecting = client.clone(); + let task = tokio::spawn(async move { connecting.connect(&url).await }); + let (mut stream, _) = tokio::time::timeout(Duration::from_secs(10), listener.accept()) + .await + .expect("client connected within 10s") + .expect("accept"); + let connect = expect_kind(&mut stream, CONNECT, MEDIUM).await; + assert_eq!( + connect.body[6], 4, + "CONNECT protocol level must be 4 (3.1.1)" + ); + stream + .write_all(&[0x20, 0x02, 0x00, 0x00]) + .await + .expect("write CONNACK"); + task.await.unwrap().expect("3.1.1 connect"); + + let properties = PublishProperties { + topic_alias: Some(1), + response_topic: Some("r/#".to_string()), + subscription_identifiers: vec![1], + ..Default::default() + }; + client + .publish_with_options("a/b", b"x".to_vec(), qos0_with(properties)) + .await + .expect("3.1.1 publish must not be failed by v5-only property checks"); + let frame = expect_kind(&mut stream, PUBLISH, MEDIUM).await; + assert_eq!( + frame.body, + vec![0x00, 0x03, b'a', b'/', b'b', b'x'], + "3.1.1 PUBLISH must carry only Topic Name and payload, no v5 properties" + ); +} diff --git a/crates/mqtt5/tests/conf_client_c.rs b/crates/mqtt5/tests/conf_client_c.rs index 1ec7774b..3b321615 100644 --- a/crates/mqtt5/tests/conf_client_c.rs +++ b/crates/mqtt5/tests/conf_client_c.rs @@ -366,6 +366,10 @@ async fn mqtt_3_14_1_1_disconnect_reserved_bits_client_sends_0x81_then_closes() packets.iter().any(|p| p.kind() == 14 && p.reason() == 0x81), "MQTT-3.14.1-1 violated: no DISCONNECT 0x81 after DISCONNECT with reserved flags 0x1; saw {seen:?}, socket closed={eof}" ); + assert!( + eof, + "MQTT-3.14.4-2 violated: network connection not closed after DISCONNECT 0x81; saw {seen:?}" + ); })) .await; } diff --git a/crates/mqtt5/tests/conf_client_d.rs b/crates/mqtt5/tests/conf_client_d.rs index b4438c9e..c1ad0f0f 100644 --- a/crates/mqtt5/tests/conf_client_d.rs +++ b/crates/mqtt5/tests/conf_client_d.rs @@ -945,8 +945,8 @@ async fn mqtt_4_9_0_1_send_quota_reinitialized_after_ack_timeout_and_reconnect() ) .await; assert!( - matches!(timed_out, Ok(Err(mqtt5::MqttError::Timeout))), - "setup: unacked QoS1 publish must hit the client ack timeout: {timed_out:?}" + matches!(&timed_out, Ok(Ok(mqtt5::PublishResult::Queued(handle))) if handle.try_outcome().is_none()), + "setup: unacked QoS1 publish must hit the client ack wait and stay in flight: {timed_out:?}" ); drop(peer); assert!( diff --git a/crates/mqtt5/tests/conf_client_offline_queue.rs b/crates/mqtt5/tests/conf_client_offline_queue.rs new file mode 100644 index 00000000..69fffa28 --- /dev/null +++ b/crates/mqtt5/tests/conf_client_offline_queue.rs @@ -0,0 +1,973 @@ +use mqtt5::{ + ConnectOptions, Delivery, IndeterminateReason, MqttClient, MqttError, PublishHandle, + PublishOptions, PublishOutcome, PublishRejection, PublishResult, QoS, +}; +use std::net::SocketAddr; +use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; +use tokio::time::{timeout, Instant}; + +const CONNECT: u8 = 1; +const PUBLISH: u8 = 3; +const PUBREL: u8 = 6; +const PINGREQ: u8 = 12; + +const PROP_RECEIVE_MAXIMUM: u8 = 0x21; +const PROP_MAXIMUM_QOS: u8 = 0x24; +const PROP_RETAIN_AVAILABLE: u8 = 0x25; +const PROP_MAXIMUM_PACKET_SIZE: u8 = 0x27; + +const W: Duration = Duration::from_secs(3); + +struct Pkt { + kind: u8, + flags: u8, + body: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct Seen { + topic: String, + qos: u8, + dup: bool, + retain: bool, + pid: Option, +} + +fn varint(mut n: usize) -> Vec { + let mut out = Vec::new(); + loop { + let mut byte = u8::try_from(n % 128).unwrap_or(0); + n /= 128; + if n > 0 { + byte |= 0x80; + } + out.push(byte); + if n == 0 { + return out; + } + } +} + +fn frame(first: u8, body: &[u8]) -> Vec { + let mut out = vec![first]; + out.extend(varint(body.len())); + out.extend_from_slice(body); + out +} + +fn connack(session_present: bool, props: &[u8]) -> Vec { + let mut body = vec![u8::from(session_present), 0]; + body.extend(varint(props.len())); + body.extend_from_slice(props); + frame(0x20, &body) +} + +fn prop_byte(id: u8, v: u8) -> Vec { + vec![id, v] +} + +fn prop_u16(id: u8, v: u16) -> Vec { + let mut out = vec![id]; + out.extend(v.to_be_bytes()); + out +} + +fn prop_u32(id: u8, v: u32) -> Vec { + let mut out = vec![id]; + out.extend(v.to_be_bytes()); + out +} + +fn puback(pid: u16) -> Vec { + frame(0x40, &pid.to_be_bytes()) +} + +fn pubrec(pid: u16) -> Vec { + frame(0x50, &pid.to_be_bytes()) +} + +fn pubcomp(pid: u16) -> Vec { + frame(0x70, &pid.to_be_bytes()) +} + +fn take_varint(b: &[u8], pos: &mut usize) -> Option { + let mut value = 0usize; + let mut mult = 1usize; + loop { + let byte = *b.get(*pos)?; + *pos += 1; + value += usize::from(byte & 0x7f) * mult; + if byte & 0x80 == 0 { + return Some(value); + } + mult *= 128; + if mult > 128 * 128 * 128 { + return None; + } + } +} + +fn take_u16(b: &[u8], pos: &mut usize) -> Option { + let v = u16::from_be_bytes([*b.get(*pos)?, *b.get(*pos + 1)?]); + *pos += 2; + Some(v) +} + +fn parse_publish(p: &Pkt) -> Option { + let b = &p.body; + let mut pos = 0; + let len = usize::from(take_u16(b, &mut pos)?); + let topic = String::from_utf8_lossy(b.get(pos..pos + len)?).into_owned(); + pos += len; + let qos = (p.flags >> 1) & 0x03; + let pid = if qos > 0 { + Some(take_u16(b, &mut pos)?) + } else { + None + }; + Some(Seen { + topic, + qos, + dup: p.flags & 0x08 != 0, + retain: p.flags & 0x01 != 0, + pid, + }) +} + +fn packet_id(p: &Pkt) -> Option { + let mut pos = 0; + take_u16(&p.body, &mut pos) +} + +fn try_split_packet(buf: &mut Vec) -> Option { + let first = *buf.first()?; + let mut pos = 1; + let len = take_varint(buf, &mut pos)?; + if buf.len() < pos + len { + return None; + } + let body = buf[pos..pos + len].to_vec(); + buf.drain(..pos + len); + Some(Pkt { + kind: first >> 4, + flags: first & 0x0f, + body, + }) +} + +struct Peer { + stream: TcpStream, + buf: Vec, +} + +impl Peer { + async fn read(&mut self, wait: Duration) -> Option { + let deadline = Instant::now() + wait; + loop { + if let Some(p) = try_split_packet(&mut self.buf) { + return Some(p); + } + let mut chunk = [0u8; 4096]; + match tokio::time::timeout_at(deadline, self.stream.read(&mut chunk)).await { + Err(_) | Ok(Ok(0) | Err(_)) => return None, + Ok(Ok(n)) => self.buf.extend_from_slice(&chunk[..n]), + } + } + } + + async fn send(&mut self, bytes: &[u8]) { + let _ = self.stream.write_all(bytes).await; + let _ = self.stream.flush().await; + } +} + +#[derive(Clone, Copy, PartialEq, Eq)] +enum Acks { + All, + None, +} + +async fn observe(peer: &mut Peer, window: Duration, acks: Acks) -> Vec { + let deadline = Instant::now() + window; + let mut seen = Vec::new(); + loop { + let remaining = deadline.saturating_duration_since(Instant::now()); + let Some(p) = peer.read(remaining).await else { + return seen; + }; + match p.kind { + PUBLISH => { + let Some(info) = parse_publish(&p) else { + continue; + }; + match (info.qos, info.pid, acks) { + (1, Some(pid), Acks::All) => peer.send(&puback(pid)).await, + (2, Some(pid), Acks::All) => { + peer.send(&pubrec(pid)).await; + } + _ => {} + } + seen.push(info); + } + PUBREL if acks == Acks::All => { + if let Some(pid) = packet_id(&p) { + peer.send(&pubcomp(pid)).await; + } + } + PINGREQ => peer.send(&[0xD0, 0x00]).await, + _ => {} + } + } +} + +async fn listener() -> (TcpListener, SocketAddr) { + let l = TcpListener::bind("127.0.0.1:0").await.expect("bind"); + let a = l.local_addr().expect("addr"); + (l, a) +} + +async fn accept_connect(l: &TcpListener) -> Option { + let (stream, _) = timeout(W, l.accept()).await.ok()?.ok()?; + let mut peer = Peer { + stream, + buf: Vec::new(), + }; + let p = peer.read(W).await?; + (p.kind == CONNECT).then_some(peer) +} + +fn persistent_opts(id: &str) -> ConnectOptions { + ConnectOptions::new(id) + .with_clean_start(false) + .with_session_expiry_interval(3600) + .with_automatic_reconnect(false) +} + +async fn connect(client: &MqttClient, l: &TcpListener, addr: SocketAddr, reply: &[u8]) -> Peer { + let pending = { + let c = client.clone(); + tokio::spawn(async move { c.connect(&format!("mqtt://{addr}")).await }) + }; + let mut peer = accept_connect(l).await.expect("CONNECT"); + peer.send(reply).await; + pending.await.expect("join").expect("connect"); + peer +} + +async fn lose(client: &MqttClient, peer: Peer) { + drop(peer); + let deadline = Instant::now() + W; + while Instant::now() < deadline { + if !client.is_connected().await { + return; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + panic!("setup: client did not notice the connection loss"); +} + +fn qos_opts(qos: QoS, retain: bool) -> PublishOptions { + PublishOptions { + qos, + retain, + ..Default::default() + } +} + +async fn queue( + client: &MqttClient, + topic: &str, + payload: Vec, + qos: QoS, + retain: bool, +) -> PublishHandle { + match client + .publish_with_options(topic, payload, qos_opts(qos, retain)) + .await + { + Ok(PublishResult::Queued(handle)) => handle, + other => panic!("setup: offline publish to {topic} must be queued, got {other:?}"), + } +} + +async fn settled(handle: PublishHandle) -> PublishOutcome { + timeout(W, handle.outcome()) + .await + .expect("publish outcome must settle") +} + +fn topics(seen: &[Seen]) -> Vec<&str> { + seen.iter().map(|s| s.topic.as_str()).collect() +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn offline_retained_publish_rejected_at_enqueue_when_retain_unavailable() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-ra-enqueue")); + let peer = connect( + &client, + &l, + addr, + &connack(false, &prop_byte(PROP_RETAIN_AVAILABLE, 0)), + ) + .await; + lose(&client, peer).await; + + let retained = client + .publish_with_options( + "oq/retained", + b"r".to_vec(), + qos_opts(QoS::AtLeastOnce, true), + ) + .await; + assert!( + matches!(retained, Err(MqttError::RetainNotSupported)), + "an offline RETAIN publish must be rejected at once when the last CONNACK said Retain Available 0, got {retained:?}" + ); + let plain = client + .publish_with_options("oq/plain", b"p".to_vec(), qos_opts(QoS::AtLeastOnce, false)) + .await; + assert!( + matches!(plain, Ok(PublishResult::Queued(_))), + "a non-retained offline publish must still be queued, got {plain:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn flush_rejects_retained_message_after_retain_becomes_unavailable() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-ra-flush")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let retained = queue( + &client, + "oq/retained", + b"r".to_vec(), + QoS::AtLeastOnce, + true, + ) + .await; + let next = queue(&client, "oq/next", b"n".to_vec(), QoS::AtLeastOnce, false).await; + + let mut peer = connect( + &client, + &l, + addr, + &connack(true, &prop_byte(PROP_RETAIN_AVAILABLE, 0)), + ) + .await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert_eq!( + topics(&seen), + vec!["oq/next"], + "the retained message must not be sent with Retain Available 0; the next one must be" + ); + assert!(seen.iter().all(|s| !s.retain)); + + assert_eq!( + settled(retained.clone()).await, + PublishOutcome::Rejected(PublishRejection::RetainNotSupported) + ); + assert!(matches!( + settled(next).await, + PublishOutcome::Delivered(Delivery::AtLeastOnce { .. }) + )); + + let live = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("oq/live", b"l".to_vec()).await }) + }; + let later = observe(&mut peer, Duration::from_millis(500), Acks::All).await; + assert_eq!(topics(&later), vec!["oq/live"]); + assert!(matches!( + live.await.expect("join"), + Ok(PublishResult::Sent(Delivery::AtLeastOnce { .. })) + )); + assert_eq!( + retained.try_outcome(), + Some(PublishOutcome::Rejected( + PublishRejection::RetainNotSupported + )), + "a settled outcome must not be changed by later acknowledgements of reused identifiers" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn flush_rejects_oversized_message_after_maximum_packet_size_shrinks() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-mps-flush")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let big = queue(&client, "oq/big", vec![0u8; 300], QoS::AtLeastOnce, false).await; + let small = queue(&client, "oq/small", b"s".to_vec(), QoS::AtLeastOnce, false).await; + + let mut peer = connect( + &client, + &l, + addr, + &connack(true, &prop_u32(PROP_MAXIMUM_PACKET_SIZE, 100)), + ) + .await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert_eq!( + topics(&seen), + vec!["oq/small"], + "the oversized message must not be sent; the next one must be" + ); + assert_eq!( + settled(big).await, + PublishOutcome::Rejected(PublishRejection::PacketTooLarge) + ); + assert!(matches!( + settled(small).await, + PublishOutcome::Delivered(Delivery::AtLeastOnce { .. }) + )); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn flush_downgrades_to_maximum_qos_and_reports_qos_used() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-maxqos-flush")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let exactly_once = queue(&client, "oq/two", b"2".to_vec(), QoS::ExactlyOnce, false).await; + + let mut peer = connect( + &client, + &l, + addr, + &connack(true, &prop_byte(PROP_MAXIMUM_QOS, 1)), + ) + .await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert_eq!(topics(&seen), vec!["oq/two"]); + assert_eq!( + seen[0].qos, 1, + "the queued QoS 2 message must go out at the new Maximum QoS 1" + ); + + let outcome = settled(exactly_once).await; + let PublishOutcome::Delivered(delivery) = outcome else { + panic!("downgraded message must be delivered, got {outcome:?}"); + }; + assert_eq!(delivery.qos_used(), QoS::AtLeastOnce); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn flush_downgrade_to_qos0_resolves_unconfirmed() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-maxqos0-flush")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let at_least_once = queue(&client, "oq/one", b"1".to_vec(), QoS::AtLeastOnce, false).await; + + let mut peer = connect( + &client, + &l, + addr, + &connack(true, &prop_byte(PROP_MAXIMUM_QOS, 0)), + ) + .await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::None).await; + assert_eq!(topics(&seen), vec!["oq/one"]); + assert_eq!(seen[0].qos, 0); + assert_eq!( + settled(at_least_once).await, + PublishOutcome::Delivered(Delivery::Unconfirmed), + "a message downgraded to QoS 0 settles as written but unconfirmed" + ); +} + +async fn flush_unacknowledged( + client: &MqttClient, + l: &TcpListener, + addr: SocketAddr, + acks: Acks, + expected: usize, +) -> Vec { + let mut peer = connect(client, l, addr, &connack(true, &[])).await; + let seen = observe(&mut peer, Duration::from_millis(800), acks).await; + assert_eq!( + seen.len(), + expected, + "setup: every queued message must be flushed once: {seen:?}" + ); + lose(client, peer).await; + seen +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn replay_skips_unacked_publishes_that_no_longer_conform() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-replay-qos1")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let retained = queue( + &client, + "oq/retained", + b"r".to_vec(), + QoS::AtLeastOnce, + true, + ) + .await; + let big = queue(&client, "oq/big", vec![0u8; 300], QoS::AtLeastOnce, false).await; + let small = queue(&client, "oq/small", b"s".to_vec(), QoS::AtLeastOnce, false).await; + flush_unacknowledged(&client, &l, addr, Acks::None, 3).await; + + let mut caps = prop_byte(PROP_RETAIN_AVAILABLE, 0); + caps.extend(prop_u32(PROP_MAXIMUM_PACKET_SIZE, 100)); + let mut peer = connect(&client, &l, addr, &connack(true, &caps)).await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert_eq!( + topics(&seen), + vec!["oq/small"], + "only the unacknowledged PUBLISH that still conforms may be replayed: {seen:?}" + ); + assert!(seen[0].dup, "the replayed PUBLISH must carry DUP=1"); + + assert_eq!( + settled(retained).await, + PublishOutcome::Indeterminate(IndeterminateReason::ReplayNotConforming) + ); + assert_eq!( + settled(big).await, + PublishOutcome::Indeterminate(IndeterminateReason::ReplayNotConforming) + ); + assert!(matches!( + settled(small).await, + PublishOutcome::Delivered(Delivery::AtLeastOnce { .. }) + )); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn replay_skips_unacked_qos2_after_maximum_qos_drops() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-replay-qos2")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let exactly_once = queue(&client, "oq/two", b"2".to_vec(), QoS::ExactlyOnce, false).await; + let first = flush_unacknowledged(&client, &l, addr, Acks::None, 1).await; + assert_eq!(first[0].qos, 2); + + let mut peer = connect( + &client, + &l, + addr, + &connack(true, &prop_byte(PROP_MAXIMUM_QOS, 1)), + ) + .await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert!( + seen.iter().all(|s| s.qos <= 1), + "a QoS 2 PUBLISH must not be replayed after the server lowered Maximum QoS to 1: {seen:?}" + ); + assert_eq!( + settled(exactly_once).await, + PublishOutcome::Indeterminate(IndeterminateReason::ReplayNotConforming) + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn session_lost_requeues_unacked_qos1_in_order() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-sp0-qos1")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let first = queue(&client, "oq/1", b"1".to_vec(), QoS::AtLeastOnce, false).await; + let second = queue(&client, "oq/2", b"2".to_vec(), QoS::AtLeastOnce, false).await; + flush_unacknowledged(&client, &l, addr, Acks::None, 2).await; + let third = queue(&client, "oq/3", b"3".to_vec(), QoS::AtLeastOnce, false).await; + + let mut peer = connect(&client, &l, addr, &connack(false, &[])).await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert_eq!( + topics(&seen), + vec!["oq/1", "oq/2", "oq/3"], + "after Session Present 0 the unacknowledged QoS 1 messages must be re-sent first, in order" + ); + assert!( + seen.iter().all(|s| !s.dup), + "messages re-sent on a new session are new PUBLISH packets (DUP=0): {seen:?}" + ); + for handle in [first, second, third] { + assert!(matches!( + settled(handle).await, + PublishOutcome::Delivered(Delivery::AtLeastOnce { .. }) + )); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn session_lost_resolves_unacked_qos2_by_exchange_stage() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-sp0-qos2")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let unreceived = queue( + &client, + "oq/publish", + b"p".to_vec(), + QoS::ExactlyOnce, + false, + ) + .await; + let released = queue(&client, "oq/pubrel", b"r".to_vec(), QoS::ExactlyOnce, false).await; + let mut peer = connect(&client, &l, addr, &connack(true, &[])).await; + let mut seen = Vec::new(); + let deadline = Instant::now() + Duration::from_millis(800); + while let Some(p) = peer + .read(deadline.saturating_duration_since(Instant::now())) + .await + { + if let Some(info) = (p.kind == PUBLISH).then(|| parse_publish(&p)).flatten() { + if let (Some(pid), "oq/pubrel") = (info.pid, info.topic.as_str()) { + peer.send(&pubrec(pid)).await; + } + seen.push(info); + } + } + assert_eq!(seen.len(), 2, "setup: both QoS 2 messages must be flushed"); + lose(&client, peer).await; + + let mut peer = connect(&client, &l, addr, &connack(false, &[])).await; + let replayed = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert!( + replayed.is_empty(), + "unacknowledged QoS 2 messages must not be re-sent on a new session: {replayed:?}" + ); + assert_eq!( + settled(unreceived).await, + PublishOutcome::Indeterminate(IndeterminateReason::SessionLost) + ); + assert!( + matches!( + settled(released).await, + PublishOutcome::Delivered(Delivery::ExactlyOnce { .. }) + ), + "after PUBREC Success the server owns the message; losing the session does not make it indeterminate" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn clean_start_discards_unacked_session_state_but_keeps_offline_queue() { + let (l, addr) = listener().await; + let client = MqttClient::with_options( + ConnectOptions::new("oq-clean-start") + .with_clean_start(true) + .with_automatic_reconnect(false), + ); + client.set_queue_on_disconnect(true).await; + + let sent = queue(&client, "oq/sent", b"s".to_vec(), QoS::AtLeastOnce, false).await; + let mut peer = connect(&client, &l, addr, &connack(false, &[])).await; + let first = observe(&mut peer, Duration::from_millis(800), Acks::None).await; + assert_eq!( + topics(&first), + vec!["oq/sent"], + "setup: queued message flushed" + ); + lose(&client, peer).await; + + let unsent = queue(&client, "oq/unsent", b"u".to_vec(), QoS::AtLeastOnce, false).await; + let mut peer = connect(&client, &l, addr, &connack(false, &[])).await; + let seen = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert_eq!( + topics(&seen), + vec!["oq/unsent"], + "Clean Start 1 discards unacknowledged session state [MQTT-3.1.2-4]; never-sent queued messages still go out" + ); + assert_eq!( + settled(sent).await, + PublishOutcome::Indeterminate(IndeterminateReason::SessionDiscarded) + ); + assert!(matches!( + settled(unsent).await, + PublishOutcome::Delivered(Delivery::AtLeastOnce { .. }) + )); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn flush_respects_receive_maximum_and_settles_each_publish() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-rm-flush")); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let mut handles = Vec::new(); + for i in 0..3u8 { + handles.push( + queue( + &client, + &format!("oq/{i}"), + vec![i], + QoS::AtLeastOnce, + false, + ) + .await, + ); + } + let mut peer = connect( + &client, + &l, + addr, + &connack(true, &prop_u16(PROP_RECEIVE_MAXIMUM, 1)), + ) + .await; + let unacked = observe(&mut peer, Duration::from_millis(500), Acks::None).await; + assert_eq!( + topics(&unacked), + vec!["oq/0"], + "only Receive Maximum 1 unacknowledged PUBLISH may be in flight" + ); + assert_eq!(handles[0].try_outcome(), None); + if let Some(pid) = unacked[0].pid { + peer.send(&puback(pid)).await; + } + let rest = observe(&mut peer, Duration::from_millis(800), Acks::All).await; + assert_eq!(topics(&rest), vec!["oq/1", "oq/2"]); + for handle in handles { + assert!(matches!( + settled(handle).await, + PublishOutcome::Delivered(Delivery::AtLeastOnce { .. }) + )); + } +} + +async fn dropped_client_settles_every_handle_abandoned(receive_maximum: Option, id: &str) { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts(id)); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + + let mut handles = Vec::new(); + for i in 0..3u8 { + handles.push( + queue( + &client, + &format!("oq/{i}"), + vec![i], + QoS::AtLeastOnce, + false, + ) + .await, + ); + } + let caps = receive_maximum.map_or_else(Vec::new, |rm| prop_u16(PROP_RECEIVE_MAXIMUM, rm)); + let mut peer = connect(&client, &l, addr, &connack(true, &caps)).await; + let unacked = observe(&mut peer, Duration::from_millis(300), Acks::None).await; + let expected = if receive_maximum.is_some() { 1 } else { 3 }; + assert_eq!( + unacked.len(), + expected, + "setup: flushed within Receive Maximum" + ); + lose(&client, peer).await; + let _ = client.disconnect().await; + drop(client); + + for (i, handle) in handles.into_iter().enumerate() { + let got = timeout(Duration::from_secs(2), handle.outcome()).await; + assert_eq!( + got.ok(), + Some(PublishOutcome::Indeterminate( + IndeterminateReason::Abandoned + )), + "handle {i} must settle Abandoned once the client is dropped" + ); + } +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn dropped_client_after_loss_mid_flush_abandons_every_handle() { + dropped_client_settles_every_handle_abandoned(Some(1), "oq-drop-mid-flush").await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn dropped_client_after_completed_flush_abandons_every_handle() { + dropped_client_settles_every_handle_abandoned(None, "oq-drop-flushed").await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn quota_waiter_fails_fast_when_the_connection_is_lost() { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts("oq-quota-loss")); + let mut peer = connect( + &client, + &l, + addr, + &connack(false, &prop_u16(PROP_RECEIVE_MAXIMUM, 1)), + ) + .await; + let first = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("live/first", b"1".to_vec()).await }) + }; + let seen = observe(&mut peer, Duration::from_millis(300), Acks::None).await; + assert_eq!(topics(&seen), vec!["live/first"]); + let waiting = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("live/second", b"2".to_vec()).await }) + }; + tokio::time::sleep(Duration::from_millis(100)).await; + lose(&client, peer).await; + + let second = timeout(W, waiting).await; + assert!( + matches!(second, Ok(Ok(Err(MqttError::NotConnected)))), + "a publish waiting for send quota must fail with NotConnected when the connection is lost, got {second:?}" + ); + assert!(matches!( + timeout(W, first).await, + Ok(Ok(Ok(PublishResult::Queued(_)))) + )); +} + +async fn live_publish_in_flight_across_loss(session_present: bool, id: &str) -> Vec { + let (l, addr) = listener().await; + let client = MqttClient::with_options(persistent_opts(id)); + let mut peer = connect(&client, &l, addr, &connack(false, &[])).await; + + let publishing = { + let c = client.clone(); + tokio::spawn(async move { + c.publish_with_options("live/1", b"x".to_vec(), qos_opts(QoS::AtLeastOnce, false)) + .await + }) + }; + let first = observe(&mut peer, Duration::from_millis(300), Acks::None).await; + assert_eq!(topics(&first), vec!["live/1"]); + lose(&client, peer).await; + let result = timeout(W, publishing) + .await + .expect("publish returns") + .expect("join"); + let Ok(PublishResult::Queued(handle)) = result else { + panic!("a publish in flight across a connection loss must return a handle, got {result:?}"); + }; + assert_eq!(handle.try_outcome(), None); + + let mut peer = connect(&client, &l, addr, &connack(session_present, &[])).await; + let resent = observe(&mut peer, Duration::from_millis(500), Acks::All).await; + assert_eq!(topics(&resent), vec!["live/1"]); + assert!(matches!( + settled(handle).await, + PublishOutcome::Delivered(Delivery::AtLeastOnce { .. }) + )); + resent +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn live_publish_in_flight_across_loss_settles_after_resume() { + let resent = live_publish_in_flight_across_loss(true, "oq-live-resume").await; + assert!(resent[0].dup, "resent on the resumed session with DUP=1"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn live_publish_in_flight_across_loss_settles_after_session_lost() { + let resent = live_publish_in_flight_across_loss(false, "oq-live-sp0").await; + assert!( + !resent[0].dup, + "re-sent as a new PUBLISH on the new session" + ); +} + +async fn hung_broker(peer: &mut Peer, client: &MqttClient) { + let deadline = Instant::now() + Duration::from_secs(8); + while Instant::now() < deadline { + let _ = peer.read(Duration::from_millis(50)).await; + if !client.is_connected().await { + return; + } + } + panic!("setup: keepalive did not detect the hung broker"); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn keepalive_loss_fails_quota_waiters_fast() { + let (l, addr) = listener().await; + let client = MqttClient::with_options( + persistent_opts("oq-ka-quota").with_keep_alive(Duration::from_secs(1)), + ); + let mut peer = connect( + &client, + &l, + addr, + &connack(false, &prop_u16(PROP_RECEIVE_MAXIMUM, 1)), + ) + .await; + let first = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("live/first", b"1".to_vec()).await }) + }; + tokio::time::sleep(Duration::from_millis(100)).await; + let waiting = { + let c = client.clone(); + tokio::spawn(async move { c.publish_qos1("live/second", b"2".to_vec()).await }) + }; + tokio::time::sleep(Duration::from_millis(100)).await; + hung_broker(&mut peer, &client).await; + let second = timeout(W, waiting).await; + let first = timeout(W, first).await; + assert!( + matches!(second, Ok(Ok(Err(MqttError::NotConnected)))), + "a quota waiter must fail fast after keepalive-detected loss: {second:?}" + ); + assert!( + matches!(first, Ok(Ok(Ok(PublishResult::Queued(_))))), + "the in-flight publish must return its handle promptly after keepalive-detected loss: {first:?}" + ); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn keepalive_loss_then_drop_abandons_handles() { + let (l, addr) = listener().await; + let client = MqttClient::with_options( + persistent_opts("oq-ka-drop").with_keep_alive(Duration::from_secs(1)), + ); + let peer = connect(&client, &l, addr, &connack(false, &[])).await; + lose(&client, peer).await; + let mut handles = Vec::new(); + for i in 0..3u8 { + handles.push( + queue( + &client, + &format!("oq/{i}"), + vec![i], + QoS::AtLeastOnce, + false, + ) + .await, + ); + } + let mut peer = connect( + &client, + &l, + addr, + &connack(true, &prop_u16(PROP_RECEIVE_MAXIMUM, 1)), + ) + .await; + hung_broker(&mut peer, &client).await; + drop(client); + for (i, handle) in handles.into_iter().enumerate() { + let got = timeout(Duration::from_secs(2), handle.outcome()).await; + let got = got.ok(); + assert_eq!( + got, + Some(PublishOutcome::Indeterminate( + IndeterminateReason::Abandoned + )), + "handle {i} after keepalive-detected loss and drop: {got:?}" + ); + } + drop(peer); +} diff --git a/crates/mqtt5/tests/conf_client_quic.rs b/crates/mqtt5/tests/conf_client_quic.rs new file mode 100644 index 00000000..c1fcd9cf --- /dev/null +++ b/crates/mqtt5/tests/conf_client_quic.rs @@ -0,0 +1,360 @@ +#![cfg(all(feature = "broker", feature = "transport-quic"))] + +use mqtt5::broker::quic_acceptor::QuicAcceptorConfig; +use mqtt5::{ConnectOptions, Message, MqttClient, QoS, SubscribeOptions}; +use std::path::Path; +use std::time::Duration; +use tokio::sync::mpsc; + +const SUBSCRIBE: u8 = 0x82; +const PUBACK: u8 = 0x40; +const DISCONNECT: u8 = 0xE0; +const UNSPECIFIED_CLOSE: u32 = 0xB2; + +struct FakeBroker { + _endpoint: quinn::Endpoint, + conn: quinn::Connection, + ctl_send: quinn::SendStream, + ctl_recv: quinn::RecvStream, +} + +async fn read_frame(recv: &mut quinn::RecvStream) -> Option> { + let mut first = [0u8; 1]; + recv.read_exact(&mut first).await.ok()?; + let mut frame = vec![first[0]]; + let mut remaining = 0usize; + let mut multiplier = 1usize; + loop { + let mut byte = [0u8; 1]; + recv.read_exact(&mut byte).await.ok()?; + frame.push(byte[0]); + remaining += usize::from(byte[0] & 0x7F) * multiplier; + multiplier *= 128; + if byte[0] & 0x80 == 0 { + break; + } + } + let mut body = vec![0u8; remaining]; + recv.read_exact(&mut body).await.ok()?; + frame.extend_from_slice(&body); + Some(frame) +} + +fn encode_varint(mut value: usize, out: &mut Vec) { + loop { + let mut byte = u8::try_from(value % 128).expect("remainder fits u8"); + value /= 128; + if value > 0 { + byte |= 0x80; + } + out.push(byte); + if value == 0 { + break; + } + } +} + +fn publish_frame( + qos: u8, + topic: &str, + packet_id: Option, + props: &[u8], + payload: &[u8], +) -> Vec { + let mut body = u16::try_from(topic.len()) + .expect("topic fits u16") + .to_be_bytes() + .to_vec(); + body.extend_from_slice(topic.as_bytes()); + if let Some(id) = packet_id { + body.extend_from_slice(&id.to_be_bytes()); + } + encode_varint(props.len(), &mut body); + body.extend_from_slice(props); + body.extend_from_slice(payload); + let mut frame = vec![0x30 | (qos << 1)]; + encode_varint(body.len(), &mut frame); + frame.extend_from_slice(&body); + frame +} + +fn topic_alias_prop(alias: u16) -> Vec { + let mut out = vec![0x23]; + out.extend_from_slice(&alias.to_be_bytes()); + out +} + +async fn connect_fake_broker(options: ConnectOptions) -> (MqttClient, FakeBroker) { + let _ = rustls::crypto::ring::default_provider().install_default(); + let cert_dir = Path::new(env!("CARGO_MANIFEST_DIR")).join("../../test_certs"); + let certs = QuicAcceptorConfig::load_cert_chain_from_file(cert_dir.join("server.pem")) + .await + .expect("server cert"); + let key = QuicAcceptorConfig::load_private_key_from_file(cert_dir.join("server.key")) + .await + .expect("server key"); + let server_config = QuicAcceptorConfig::new(certs, key) + .with_alpn_protocols(vec![b"mqtt".to_vec()]) + .build_server_config() + .expect("server config"); + let endpoint = quinn::Endpoint::server(server_config, "127.0.0.1:0".parse().unwrap()) + .expect("server endpoint"); + let url = format!("quic://{}", endpoint.local_addr().expect("local addr")); + + let client = MqttClient::with_options(options.clone()); + client.set_insecure_tls(true).await; + let connecting = client.clone(); + let connect = + tokio::spawn(async move { Box::pin(connecting.connect_with_options(&url, options)).await }); + + let conn = endpoint + .accept() + .await + .expect("incoming connection") + .await + .expect("handshake"); + let (mut ctl_send, mut ctl_recv) = conn.accept_bi().await.expect("control stream"); + let connect_frame = read_frame(&mut ctl_recv).await.expect("CONNECT"); + assert_eq!(connect_frame[0] >> 4, 1, "first packet must be CONNECT"); + ctl_send + .write_all(&[0x20, 0x03, 0x00, 0x00, 0x00]) + .await + .expect("write CONNACK"); + connect + .await + .expect("connect task") + .expect("client connects to fake broker"); + ( + client, + FakeBroker { + _endpoint: endpoint, + conn, + ctl_send, + ctl_recv, + }, + ) +} + +async fn subscribe( + client: &MqttClient, + broker: &mut FakeBroker, + filter: &str, + qos: QoS, +) -> mpsc::UnboundedReceiver { + let (tx, rx) = mpsc::unbounded_channel(); + let subscriber = client.clone(); + let filter = filter.to_string(); + let options = SubscribeOptions { + qos, + ..Default::default() + }; + let task = tokio::spawn(async move { + subscriber + .subscribe_with_options(filter, options, move |message| { + let _ = tx.send(message); + }) + .await + }); + let frame = read_frame(&mut broker.ctl_recv).await.expect("SUBSCRIBE"); + assert_eq!(frame[0], SUBSCRIBE); + broker + .ctl_send + .write_all(&[0x90, 0x04, frame[2], frame[3], 0x00, qos as u8]) + .await + .expect("write SUBACK"); + task.await.expect("subscribe task").expect("subscribe"); + rx +} + +async fn drain(rx: &mut mpsc::UnboundedReceiver, wait: Duration) -> Vec { + let mut messages = Vec::new(); + while let Ok(Some(message)) = tokio::time::timeout(wait, rx.recv()).await { + messages.push(message); + } + messages +} + +async fn expect_disconnect_and_close(broker: &mut FakeBroker, reason: u8) { + let frame = tokio::time::timeout(Duration::from_secs(3), read_frame(&mut broker.ctl_recv)) + .await + .expect("client answers within 3s"); + let frame = frame.expect("DISCONNECT on the control stream"); + assert_eq!(frame[0], DISCONNECT, "client must send DISCONNECT"); + assert_eq!(frame[2], reason, "DISCONNECT reason code"); + let closed = tokio::time::timeout(Duration::from_secs(3), broker.conn.closed()).await; + match closed { + Ok(quinn::ConnectionError::ApplicationClosed(close)) => assert_eq!( + close.error_code.into_inner(), + u64::from(UNSPECIFIED_CLOSE), + "error teardown must close with an application error code" + ), + other => panic!("[MQTT-3.14.4-2] QUIC connection not closed by the client: {other:?}"), + } +} + +#[tokio::test] +async fn mqtt_3_14_4_1_quic_protocol_error_closes_connection_and_stops_acking() { + let options = ConnectOptions::new("conf-quic-close").with_automatic_reconnect(false); + let (client, mut broker) = connect_fake_broker(options).await; + let mut rx = subscribe(&client, &mut broker, "t", QoS::AtLeastOnce).await; + + broker + .ctl_send + .write_all(&[0xD1, 0x00]) + .await + .expect("write malformed PINGRESP"); + expect_disconnect_and_close(&mut broker, 0x81).await; + assert!(!client.is_connected().await); + + let acked_after_disconnect = match broker.conn.open_bi().await { + Ok((mut data_send, mut data_recv)) => { + let _ = data_send + .write_all(&publish_frame(1, "t", Some(7), &[], b"x")) + .await; + let next = + tokio::time::timeout(Duration::from_secs(2), read_frame(&mut data_recv)).await; + matches!(next, Ok(Some(frame)) if frame[0] == PUBACK) + } + Err(_) => false, + }; + assert!( + !acked_after_disconnect, + "[MQTT-3.14.4-1] client wrote PUBACK after its own DISCONNECT" + ); + assert!( + drain(&mut rx, Duration::from_millis(300)).await.is_empty(), + "[MQTT-3.14.4-1] client delivered a PUBLISH received after its own DISCONNECT" + ); + + let _ = client.disconnect().await; + assert!( + client.quic_connection().await.is_none(), + "disconnect() after a reader-initiated close must release the QUIC connection" + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_24_quic_control_stream_enforces_client_maximum_packet_size() { + let mut options = ConnectOptions::new("conf-quic-mps-ctl").with_automatic_reconnect(false); + options.properties.maximum_packet_size = Some(64); + let (client, mut broker) = connect_fake_broker(options).await; + let mut rx = subscribe(&client, &mut broker, "t", QoS::AtMostOnce).await; + + broker + .ctl_send + .write_all(&publish_frame(0, "t", None, &[], &[b'x'; 200])) + .await + .expect("write oversized PUBLISH"); + expect_disconnect_and_close(&mut broker, 0x95).await; + assert!( + drain(&mut rx, Duration::from_millis(300)).await.is_empty(), + "[MQTT-3.1.2-24] oversized PUBLISH on the control stream was delivered" + ); +} + +#[tokio::test] +async fn mqtt_3_1_2_24_quic_data_stream_enforces_client_maximum_packet_size() { + let mut options = ConnectOptions::new("conf-quic-mps-data").with_automatic_reconnect(false); + options.properties.maximum_packet_size = Some(64); + let (client, mut broker) = connect_fake_broker(options).await; + let mut rx = subscribe(&client, &mut broker, "t", QoS::AtMostOnce).await; + + let (mut data_send, _data_recv) = broker.conn.open_bi().await.expect("open data stream"); + data_send + .write_all(&publish_frame(0, "t", None, &[], &[b'x'; 200])) + .await + .expect("write oversized PUBLISH"); + expect_disconnect_and_close(&mut broker, 0x95).await; + assert!( + drain(&mut rx, Duration::from_millis(300)).await.is_empty(), + "[MQTT-3.1.2-24] oversized PUBLISH on a data stream was delivered" + ); +} + +#[tokio::test] +async fn mqtt_3_3_2_10_quic_unidirectional_stream_resolves_topic_alias() { + let mut options = ConnectOptions::new("conf-quic-uni-alias").with_automatic_reconnect(false); + options.properties.topic_alias_maximum = Some(2); + let (client, mut broker) = connect_fake_broker(options).await; + let mut rx = subscribe(&client, &mut broker, "t/a", QoS::AtMostOnce).await; + + let mut uni = broker.conn.open_uni().await.expect("open uni stream"); + uni.write_all(&publish_frame(0, "t/a", None, &topic_alias_prop(1), b"one")) + .await + .expect("write aliased PUBLISH"); + uni.write_all(&publish_frame(0, "", None, &topic_alias_prop(1), b"two")) + .await + .expect("write alias-only PUBLISH"); + let topics: Vec = drain(&mut rx, Duration::from_millis(600)) + .await + .into_iter() + .map(|m| m.topic) + .collect(); + assert_eq!( + topics, + vec!["t/a", "t/a"], + "[MQTT-3.3.2-10] inbound Topic Alias on a QUIC unidirectional stream was not resolved" + ); +} + +async fn publish_acknowledged_on_server_stream( + level: QoS, + ack_type: u8, +) -> ( + tokio::task::JoinHandle>, + FakeBroker, +) { + let options = ConnectOptions::new("conf-quic-data-ack").with_automatic_reconnect(false); + let (client, mut broker) = connect_fake_broker(options).await; + let publisher = client.clone(); + let publishing = + tokio::spawn(async move { publisher.publish_qos("t/ack", b"x".to_vec(), level).await }); + let frame = read_frame(&mut broker.ctl_recv).await.expect("PUBLISH"); + assert_eq!( + frame[0] >> 4, + 3, + "client must send its PUBLISH on the control stream" + ); + let topic_len = usize::from(u16::from_be_bytes([frame[2], frame[3]])); + let pid = [frame[4 + topic_len], frame[5 + topic_len]]; + let (mut data_send, _data_recv) = broker.conn.open_bi().await.expect("open data stream"); + data_send + .write_all(&[ack_type, 0x02, pid[0], pid[1]]) + .await + .expect("write acknowledgement on a server-opened stream"); + (publishing, broker) +} + +#[tokio::test] +async fn quic_puback_on_server_opened_stream_settles_the_publish() { + let (publishing, _broker) = + publish_acknowledged_on_server_stream(QoS::AtLeastOnce, PUBACK).await; + let result = tokio::time::timeout(Duration::from_secs(3), publishing) + .await + .expect("publish settles within 3s") + .expect("publish task"); + assert!( + matches!( + result, + Ok(mqtt5::PublishResult::Sent( + mqtt5::Delivery::AtLeastOnce { .. } + )) + ), + "a PUBACK on a server-opened QUIC stream must settle the publish as delivered: {result:?}" + ); +} + +#[tokio::test] +async fn quic_mismatched_ack_on_server_opened_stream_is_a_protocol_error() { + let (publishing, mut broker) = + publish_acknowledged_on_server_stream(QoS::ExactlyOnce, PUBACK).await; + expect_disconnect_and_close(&mut broker, 0x82).await; + let result = tokio::time::timeout(Duration::from_secs(3), publishing) + .await + .expect("publish returns within 3s") + .expect("publish task"); + assert!( + !matches!(result, Ok(mqtt5::PublishResult::Sent(_))), + "a PUBACK for a QoS 2 publish must not be accepted as its acknowledgement: {result:?}" + ); +} diff --git a/crates/mqtt5/tests/integration_complete_flow.rs b/crates/mqtt5/tests/integration_complete_flow.rs index 9458c53c..bb6adb4b 100644 --- a/crates/mqtt5/tests/integration_complete_flow.rs +++ b/crates/mqtt5/tests/integration_complete_flow.rs @@ -4,7 +4,9 @@ mod common; use common::{create_test_client_with_broker, test_client_id, TestBroker}; use mqtt5::time::Duration; -use mqtt5::{ConnectOptions, MqttClient, PublishOptions, PublishResult, QoS, SubscribeOptions}; +use mqtt5::{ + ConnectOptions, Delivery, MqttClient, PublishOptions, PublishResult, QoS, SubscribeOptions, +}; use std::collections::HashMap; use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::Arc; @@ -74,8 +76,10 @@ async fn test_complete_mqtt_flow() { .expect("Failed to publish"); match result { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => assert!(packet_id > 0), + other => panic!("expected an acknowledged publish, got {other:?}"), } assert!( @@ -224,15 +228,17 @@ async fn test_qos_levels_and_acknowledgments() { .publish("test/qos0", b"QoS 0 message") .await .expect("Failed to publish QoS 0"); - assert!(matches!(result, PublishResult::QoS0)); + assert!(matches!(result, PublishResult::Sent(Delivery::Unconfirmed))); let result = client .publish_qos1("test/qos1", b"QoS 1 message") .await .expect("Failed to publish QoS 1"); match result { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => assert!(packet_id > 0), + other => panic!("expected an acknowledged publish, got {other:?}"), } let result = client @@ -240,8 +246,10 @@ async fn test_qos_levels_and_acknowledgments() { .await .expect("Failed to publish QoS 2"); match result { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => assert!(packet_id > 0), + other => panic!("expected an acknowledged publish, got {other:?}"), } let received_qos = Arc::new(Mutex::new(Vec::new())); diff --git a/crates/mqtt5/tests/integration_no_local.rs b/crates/mqtt5/tests/integration_no_local.rs index b4dc8290..48a8d9fe 100644 --- a/crates/mqtt5/tests/integration_no_local.rs +++ b/crates/mqtt5/tests/integration_no_local.rs @@ -1,8 +1,7 @@ #![cfg(feature = "broker")] -use mqtt5::broker::router::MessageRouter; +use mqtt5::broker::router::{MessageRouter, SubscriptionRequest}; use mqtt5::packet::publish::PublishPacket; use mqtt5::time::Duration; -use mqtt5::types::ProtocolVersion; use mqtt5::QoS; use std::sync::Arc; @@ -26,16 +25,8 @@ async fn test_no_local_true_filters_own_messages() { router .subscribe( - "test_client".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - true, - false, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("test_client", "test/topic", QoS::AtMostOnce) + .with_no_local(true), ) .await .unwrap(); @@ -71,18 +62,11 @@ async fn test_no_local_false_allows_own_messages() { .await; router - .subscribe( - "test_client".to_string(), - "test/topic".to_string(), + .subscribe(SubscriptionRequest::new( + "test_client", + "test/topic", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -133,33 +117,18 @@ async fn test_no_local_other_clients_receive_messages() { router .subscribe( - "publisher".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - true, - false, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("publisher", "test/topic", QoS::AtMostOnce) + .with_no_local(true), ) .await .unwrap(); router - .subscribe( - "subscriber".to_string(), - "test/topic".to_string(), + .subscribe(SubscriptionRequest::new( + "subscriber", + "test/topic", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -205,16 +174,7 @@ async fn test_no_local_with_wildcards() { router .subscribe( - "test_client".to_string(), - "test/+".to_string(), - QoS::AtMostOnce, - None, - true, - false, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("test_client", "test/+", QoS::AtMostOnce).with_no_local(true), ) .await .unwrap(); @@ -258,16 +218,7 @@ async fn test_no_local_with_multilevel_wildcard() { router .subscribe( - "test_client".to_string(), - "test/#".to_string(), - QoS::AtMostOnce, - None, - true, - false, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("test_client", "test/#", QoS::AtMostOnce).with_no_local(true), ) .await .unwrap(); @@ -301,16 +252,8 @@ async fn test_no_local_server_generated_messages() { router .subscribe( - "test_client".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - true, - false, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("test_client", "test/topic", QoS::AtMostOnce) + .with_no_local(true), ) .await .unwrap(); @@ -351,33 +294,18 @@ async fn test_no_local_multiple_subscriptions_same_client() { router .subscribe( - "test_client".to_string(), - "test/topic1".to_string(), - QoS::AtMostOnce, - None, - true, - false, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("test_client", "test/topic1", QoS::AtMostOnce) + .with_no_local(true), ) .await .unwrap(); router - .subscribe( - "test_client".to_string(), - "test/topic2".to_string(), + .subscribe(SubscriptionRequest::new( + "test_client", + "test/topic2", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -430,16 +358,8 @@ async fn test_no_local_with_qos_levels() { router .subscribe( - "test_client".to_string(), - "test/topic".to_string(), - QoS::AtLeastOnce, - None, - true, - false, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("test_client", "test/topic", QoS::AtLeastOnce) + .with_no_local(true), ) .await .unwrap(); diff --git a/crates/mqtt5/tests/integration_retain_as_published.rs b/crates/mqtt5/tests/integration_retain_as_published.rs index 11cb555a..d12ffce6 100644 --- a/crates/mqtt5/tests/integration_retain_as_published.rs +++ b/crates/mqtt5/tests/integration_retain_as_published.rs @@ -1,8 +1,7 @@ #![cfg(feature = "broker")] -use mqtt5::broker::router::MessageRouter; +use mqtt5::broker::router::{MessageRouter, SubscriptionRequest}; use mqtt5::packet::publish::PublishPacket; use mqtt5::time::Duration; -use mqtt5::types::ProtocolVersion; use mqtt5::QoS; use std::sync::Arc; @@ -26,18 +25,11 @@ async fn test_retain_as_published_false_clears_retain_flag() { .await; router - .subscribe( - "subscriber".to_string(), - "test/topic".to_string(), + .subscribe(SubscriptionRequest::new( + "subscriber", + "test/topic", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -75,16 +67,8 @@ async fn test_retain_as_published_true_preserves_retain_flag() { router .subscribe( - "subscriber".to_string(), - "test/topic".to_string(), - QoS::AtMostOnce, - None, - false, - true, - 0, - ProtocolVersion::V5, - false, - None, + SubscriptionRequest::new("subscriber", "test/topic", QoS::AtMostOnce) + .with_retain_as_published(true), ) .await .unwrap(); diff --git a/crates/mqtt5/tests/message_queuing.rs b/crates/mqtt5/tests/message_queuing.rs index 15b06dc5..b5a220f8 100644 --- a/crates/mqtt5/tests/message_queuing.rs +++ b/crates/mqtt5/tests/message_queuing.rs @@ -1,16 +1,12 @@ use mqtt5::{ConnectOptions, MqttClient, PublishOptions, PublishResult, QoS}; #[tokio::test] - async fn test_message_queuing_when_disconnected() { - // Create client with clean_start=false to enable queuing let options = ConnectOptions::new("test-client").with_clean_start(false); let client = MqttClient::with_options(options); - // Queuing should be enabled for persistent sessions assert!(client.is_queue_on_disconnect().await); - // Try to publish while disconnected - should queue the message let options = PublishOptions { qos: QoS::AtLeastOnce, ..Default::default() @@ -20,11 +16,10 @@ async fn test_message_queuing_when_disconnected() { .publish_with_options("test/topic", "queued message", options) .await; - // Should return a packet ID even though we're not connected assert!(result.is_ok()); match result.unwrap() { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result"), + PublishResult::Queued(handle) => assert_eq!(handle.try_outcome(), None), + other @ PublishResult::Sent(_) => panic!("expected a queued publish, got {other:?}"), } } @@ -33,10 +28,8 @@ async fn test_message_queuing_disabled() { let options = ConnectOptions::new("test-client").with_clean_start(true); let client = MqttClient::with_options(options); - // Queuing should be disabled for clean sessions assert!(!client.is_queue_on_disconnect().await); - // Try to publish while disconnected - should fail let options = PublishOptions { qos: QoS::AtLeastOnce, ..Default::default() @@ -52,25 +45,21 @@ async fn test_message_queuing_disabled() { async fn test_qos0_not_queued() { let client = MqttClient::new("test-client"); - // QoS 0 messages should not be queued - let options = PublishOptions::default(); // QoS 0 by default + let options = PublishOptions::default(); let result = client .publish_with_options("test/topic", "qos0 message", options) .await; - assert!(result.is_err()); // Should fail with NotConnected + assert!(result.is_err()); } #[tokio::test] - async fn test_queue_multiple_messages() { - // Create client with clean_start=false to enable queuing let options = ConnectOptions::new("test-client").with_clean_start(false); let client = MqttClient::with_options(options); - let mut packet_ids = Vec::new(); + let mut handles = Vec::new(); - // Queue multiple messages for i in 0..5 { let options = PublishOptions { qos: QoS::AtLeastOnce, @@ -83,33 +72,25 @@ async fn test_queue_multiple_messages() { assert!(result.is_ok()); match result.unwrap() { - PublishResult::QoS1Or2 { packet_id } => packet_ids.push(packet_id), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result"), + PublishResult::Queued(handle) => handles.push(handle), + other @ PublishResult::Sent(_) => panic!("expected a queued publish, got {other:?}"), } } - // All packet IDs should be unique - let mut unique_ids = packet_ids.clone(); - unique_ids.sort_unstable(); - unique_ids.dedup(); - assert_eq!(packet_ids.len(), unique_ids.len()); + assert_eq!(handles.len(), 5); + assert!(handles.iter().all(|handle| handle.try_outcome().is_none())); } #[tokio::test] - async fn test_toggle_queue_on_disconnect() { - // Create client with clean_start=false to enable queuing initially let options = ConnectOptions::new("test-client").with_clean_start(false); let client = MqttClient::with_options(options); - // Should be enabled for persistent sessions assert!(client.is_queue_on_disconnect().await); - // Disable queuing client.set_queue_on_disconnect(false).await; assert!(!client.is_queue_on_disconnect().await); - // Try to publish - should fail let options = PublishOptions { qos: QoS::AtLeastOnce, ..Default::default() @@ -119,11 +100,9 @@ async fn test_toggle_queue_on_disconnect() { .await; assert!(result.is_err()); - // Re-enable queuing client.set_queue_on_disconnect(true).await; assert!(client.is_queue_on_disconnect().await); - // Try to publish - should succeed let options = PublishOptions { qos: QoS::AtLeastOnce, ..Default::default() @@ -135,23 +114,17 @@ async fn test_toggle_queue_on_disconnect() { } #[tokio::test] - async fn test_message_replay_on_reconnect() { - // This test would require a mock broker to verify that messages are replayed - // For now, we just test the queueing behavior - - // Create client with clean_start=false to enable queuing let options = ConnectOptions::new("test-client").with_clean_start(false); let client = MqttClient::with_options(options); - // Queue several messages let messages = vec![ ("test/1", "message 1", QoS::AtLeastOnce), ("test/2", "message 2", QoS::ExactlyOnce), ("test/3", "message 3", QoS::AtLeastOnce), ]; - let mut packet_ids = Vec::new(); + let mut handles = Vec::new(); for (topic, payload, qos) in messages { let options = PublishOptions { @@ -162,23 +135,16 @@ async fn test_message_replay_on_reconnect() { let result = client.publish_with_options(topic, payload, options).await; assert!(result.is_ok()); match result.unwrap() { - PublishResult::QoS1Or2 { packet_id } => packet_ids.push(packet_id), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result"), + PublishResult::Queued(handle) => handles.push(handle), + other @ PublishResult::Sent(_) => panic!("expected a queued publish, got {other:?}"), } } - assert_eq!(packet_ids.len(), 3); - - // In a real test, we would: - // 1. Connect to a broker - // 2. Verify that all queued messages are sent with DUP flag - // 3. Verify that they maintain their original QoS levels + assert_eq!(handles.len(), 3); } #[tokio::test] - async fn test_retained_message_queuing() { - // Create client with clean_start=false to enable queuing let options = ConnectOptions::new("test-client").with_clean_start(false); let client = MqttClient::with_options(options); @@ -192,24 +158,18 @@ async fn test_retained_message_queuing() { .publish_with_options("test/retained", "retained message", options) .await; assert!(result.is_ok()); - - // The retained flag should be preserved when the message is replayed } #[tokio::test] - async fn test_clean_session_no_queuing() { let options = ConnectOptions::new("clean-client").with_clean_start(true); let client = MqttClient::with_options(options); - // Verify queuing is disabled for clean sessions assert!(!client.is_queue_on_disconnect().await); - // But we can manually enable it if needed client.set_queue_on_disconnect(true).await; assert!(client.is_queue_on_disconnect().await); - // Now queuing should work even for clean session let pub_opts = PublishOptions { qos: QoS::AtLeastOnce, ..Default::default() diff --git a/crates/mqtt5/tests/mock_client.rs b/crates/mqtt5/tests/mock_client.rs index 3a5459c2..56e066a2 100644 --- a/crates/mqtt5/tests/mock_client.rs +++ b/crates/mqtt5/tests/mock_client.rs @@ -1,6 +1,6 @@ //! Tests for the mock MQTT client functionality -use mqtt5::{MockCall, MockMqttClient, MqttClientTrait, PublishResult, QoS}; +use mqtt5::{Delivery, MockCall, MockMqttClient, MqttClientTrait, PublishResult, QoS}; #[tokio::test] async fn test_mock_client_creation() { @@ -15,12 +15,10 @@ async fn test_mock_client_creation() { async fn test_mock_client_connect_disconnect() { let mock = MockMqttClient::new("test-client"); - // Test connect let result = mock.connect("mqtt://localhost:1883").await; assert!(result.is_ok()); assert!(mock.is_connected().await); - // Verify call was recorded let calls = mock.get_calls().await; assert_eq!(calls.len(), 1); assert!(matches!( @@ -28,12 +26,10 @@ async fn test_mock_client_connect_disconnect() { MockCall::Connect { ref address } if address == "mqtt://localhost:1883" )); - // Test disconnect let result = mock.disconnect().await; assert!(result.is_ok()); assert!(!mock.is_connected().await); - // Verify both calls were recorded let calls = mock.get_calls().await; assert_eq!(calls.len(), 2); assert!(matches!(calls[1], MockCall::Disconnect)); @@ -44,21 +40,21 @@ async fn test_mock_client_publish() { let mock = MockMqttClient::new("test-client"); mock.set_connected(true); - // Test QoS 0 publish let result = mock.publish("test/topic", b"test message").await; assert!(result.is_ok()); - assert!(matches!(result.unwrap(), PublishResult::QoS0)); + assert!(matches!( + result.unwrap(), + PublishResult::Sent(Delivery::Unconfirmed) + )); - // Test QoS 1 publish let result = mock.publish_qos1("test/topic", b"test message").await; assert!(result.is_ok()); - if let PublishResult::QoS1Or2 { packet_id } = result.unwrap() { + if let PublishResult::Sent(Delivery::AtLeastOnce { packet_id }) = result.unwrap() { assert!(packet_id > 0); } else { - panic!("Expected QoS1Or2 result"); + panic!("expected an acknowledged QoS 1 result"); } - // Verify calls were recorded let calls = mock.get_calls().await; assert_eq!(calls.len(), 2); assert!(matches!( @@ -73,7 +69,6 @@ async fn test_mock_client_subscribe() { let mock = MockMqttClient::new("test-client"); mock.set_connected(true); - // Test subscribe let result = mock .subscribe("test/topic", |_msg| { println!("Received message in test"); @@ -85,7 +80,6 @@ async fn test_mock_client_subscribe() { assert!(packet_id > 0); assert_eq!(qos, QoS::AtMostOnce); - // Verify call was recorded let calls = mock.get_calls().await; assert_eq!(calls.len(), 1); assert!(matches!( @@ -102,7 +96,6 @@ async fn test_mock_client_simulate_message() { let received_messages = Arc::new(Mutex::new(Vec::new())); let messages_clone = Arc::clone(&received_messages); - // Subscribe to a topic let result = mock .subscribe("test/+", move |msg| { messages_clone @@ -113,13 +106,11 @@ async fn test_mock_client_simulate_message() { .await; assert!(result.is_ok()); - // Simulate a message let result = mock .simulate_message("test/hello", b"world".to_vec(), QoS::AtMostOnce) .await; assert!(result.is_ok()); - // Check that callback was called let messages = received_messages.lock().unwrap(); assert_eq!(messages.len(), 1); assert_eq!(messages[0].0, "test/hello"); @@ -130,20 +121,19 @@ async fn test_mock_client_simulate_message() { async fn test_mock_client_configured_responses() { let mock = MockMqttClient::new("test-client"); - // Configure a custom publish response - mock.set_publish_response(Ok(PublishResult::QoS1Or2 { packet_id: 42 })) - .await; + mock.set_publish_response(Ok(PublishResult::Sent(Delivery::AtLeastOnce { + packet_id: 42, + }))) + .await; - // Test that the configured response is returned let result = mock.publish("test/topic", b"test").await; assert!(result.is_ok()); - if let PublishResult::QoS1Or2 { packet_id } = result.unwrap() { + if let PublishResult::Sent(Delivery::AtLeastOnce { packet_id }) = result.unwrap() { assert_eq!(packet_id, 42); } else { - panic!("Expected configured QoS1Or2 result"); + panic!("expected the configured QoS 1 result"); } - // Configure a custom subscribe response mock.set_subscribe_response(Ok((123, QoS::ExactlyOnce))) .await; @@ -158,25 +148,21 @@ async fn test_mock_client_configured_responses() { async fn test_mock_client_call_tracking() { let mock = MockMqttClient::new("test-client"); - // Perform various operations let _ = mock.connect("mqtt://localhost:1883").await; let _ = mock.publish("topic1", b"msg1").await; let _ = mock.subscribe("topic2", |_| {}).await; let _ = mock.unsubscribe("topic2").await; let _ = mock.disconnect().await; - // Check that all calls were recorded let calls = mock.get_calls().await; assert_eq!(calls.len(), 5); - // Verify call types assert!(matches!(calls[0], MockCall::Connect { .. })); assert!(matches!(calls[1], MockCall::Publish { .. })); assert!(matches!(calls[2], MockCall::Subscribe { .. })); assert!(matches!(calls[3], MockCall::Unsubscribe { .. })); assert!(matches!(calls[4], MockCall::Disconnect)); - // Clear calls and verify mock.clear_calls().await; let calls = mock.get_calls().await; assert_eq!(calls.len(), 0); diff --git a/crates/mqtt5/tests/outbound_receive_maximum.rs b/crates/mqtt5/tests/outbound_receive_maximum.rs index dbe5f634..4432ba20 100644 --- a/crates/mqtt5/tests/outbound_receive_maximum.rs +++ b/crates/mqtt5/tests/outbound_receive_maximum.rs @@ -157,8 +157,8 @@ async fn ack_timeout_holds_quota_and_does_not_exceed_window() { let first_result = first_handle.await.unwrap(); assert!( - matches!(first_result, Err(MqttError::Timeout)), - "an unacknowledged publish must eventually time out" + matches!(&first_result, Ok(mqtt5::PublishResult::Queued(handle)) if handle.try_outcome().is_none()), + "an unacknowledged publish must stop waiting and stay in flight: {first_result:?}" ); let second = client.clone(); diff --git a/crates/mqtt5/tests/qos_flow.rs b/crates/mqtt5/tests/qos_flow.rs index a20a5a2d..faac5f8d 100644 --- a/crates/mqtt5/tests/qos_flow.rs +++ b/crates/mqtt5/tests/qos_flow.rs @@ -6,7 +6,7 @@ use common::TestBroker; use mqtt5::broker::config::{BrokerConfig, StorageBackend, StorageConfig}; use mqtt5::time::Duration; -use mqtt5::{MqttClient, PublishOptions, PublishResult, QoS}; +use mqtt5::{Delivery, MqttClient, PublishOptions, PublishResult, QoS}; use std::net::SocketAddr; use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::Arc; @@ -14,7 +14,6 @@ use tokio::time::sleep; #[tokio::test] async fn test_qos0_fire_and_forget() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("qos0-pub"); @@ -26,7 +25,6 @@ async fn test_qos0_fire_and_forget() { let received = Arc::new(AtomicU32::new(0)); let received_clone = received.clone(); - // Subscribe with QoS 0 sub_client .subscribe("test/qos0", move |_msg| { received_clone.fetch_add(1, Ordering::Relaxed); @@ -34,7 +32,6 @@ async fn test_qos0_fire_and_forget() { .await .unwrap(); - // Publish 10 messages with QoS 0 for i in 0..10 { let result = pub_client .publish_qos0("test/qos0", format!("Message {i}")) @@ -42,13 +39,11 @@ async fn test_qos0_fire_and_forget() { assert!(result.is_ok()); } - // Give some time for messages to arrive sleep(Duration::from_millis(500)).await; - // With QoS 0, we might not receive all messages let count = received.load(Ordering::Relaxed); println!("QoS 0: Received {count} of 10 messages"); - assert!(count > 0); // Should receive at least some messages + assert!(count > 0); pub_client.disconnect().await.unwrap(); sub_client.disconnect().await.unwrap(); @@ -56,7 +51,6 @@ async fn test_qos0_fire_and_forget() { #[tokio::test] async fn test_qos1_at_least_once() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("qos1-pub"); @@ -68,7 +62,6 @@ async fn test_qos1_at_least_once() { let received = Arc::new(AtomicU32::new(0)); let received_clone = received.clone(); - // Subscribe with QoS 1 sub_client .subscribe_with_options( "test/qos1", @@ -83,7 +76,6 @@ async fn test_qos1_at_least_once() { .await .unwrap(); - // Publish 10 messages with QoS 1 let mut packet_ids = Vec::new(); for i in 0..10 { let result = pub_client @@ -91,21 +83,20 @@ async fn test_qos1_at_least_once() { .await .unwrap(); match result { - PublishResult::QoS1Or2 { packet_id } => packet_ids.push(packet_id), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => packet_ids.push(packet_id), + other => panic!("expected an acknowledged publish, got {other:?}"), } } - // All packet IDs should be unique let mut unique_ids = packet_ids.clone(); unique_ids.sort_unstable(); unique_ids.dedup(); assert_eq!(packet_ids.len(), unique_ids.len()); - // Give time for acknowledgments sleep(Duration::from_millis(500)).await; - // With QoS 1, we should receive all messages let count = received.load(Ordering::Relaxed); assert_eq!(count, 10, "QoS 1: Should receive exactly 10 messages"); @@ -115,7 +106,6 @@ async fn test_qos1_at_least_once() { #[tokio::test] async fn test_qos2_exactly_once() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("qos2-pub"); @@ -127,7 +117,6 @@ async fn test_qos2_exactly_once() { let received = Arc::new(AtomicU32::new(0)); let received_clone = received.clone(); - // Subscribe with QoS 2 sub_client .subscribe_with_options( "test/qos2", @@ -142,7 +131,6 @@ async fn test_qos2_exactly_once() { .await .unwrap(); - // Publish 10 messages with QoS 2 let mut packet_ids = Vec::new(); for i in 0..10 { let result = pub_client @@ -150,21 +138,20 @@ async fn test_qos2_exactly_once() { .await .unwrap(); match result { - PublishResult::QoS1Or2 { packet_id } => packet_ids.push(packet_id), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => packet_ids.push(packet_id), + other => panic!("expected an acknowledged publish, got {other:?}"), } } - // All packet IDs should be unique let mut unique_ids = packet_ids.clone(); unique_ids.sort_unstable(); unique_ids.dedup(); assert_eq!(packet_ids.len(), unique_ids.len()); - // Give time for full QoS 2 handshake sleep(Duration::from_secs(1)).await; - // With QoS 2, we should receive exactly one copy of each message let count = received.load(Ordering::Relaxed); assert_eq!( count, 10, @@ -177,7 +164,6 @@ async fn test_qos2_exactly_once() { #[tokio::test] async fn test_qos_downgrade() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("qos-downgrade-pub"); @@ -189,7 +175,6 @@ async fn test_qos_downgrade() { let received_qos = Arc::new(AtomicU32::new(0)); let received_qos_clone = received_qos.clone(); - // Subscribe with QoS 0 sub_client .subscribe_with_options( "test/downgrade", @@ -198,14 +183,12 @@ async fn test_qos_downgrade() { ..Default::default() }, move |msg| { - // Store the received QoS level received_qos_clone.store(msg.qos as u32, Ordering::Relaxed); }, ) .await .unwrap(); - // Publish with QoS 2 pub_client .publish_qos2("test/downgrade", "Test message") .await @@ -213,7 +196,6 @@ async fn test_qos_downgrade() { sleep(Duration::from_millis(500)).await; - // Message should be received with downgraded QoS 0 let final_qos = received_qos.load(Ordering::Relaxed); assert_eq!(final_qos, 0, "Message should be downgraded to QoS 0"); @@ -223,7 +205,6 @@ async fn test_qos_downgrade() { #[tokio::test] async fn test_qos_upgrade_not_allowed() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("qos-upgrade-pub"); @@ -232,10 +213,9 @@ async fn test_qos_upgrade_not_allowed() { pub_client.connect(broker.address()).await.unwrap(); sub_client.connect(broker.address()).await.unwrap(); - let received_qos = Arc::new(AtomicU32::new(3)); // Invalid initial value + let received_qos = Arc::new(AtomicU32::new(3)); let received_qos_clone = received_qos.clone(); - // Subscribe with QoS 2 sub_client .subscribe_with_options( "test/upgrade", @@ -250,7 +230,6 @@ async fn test_qos_upgrade_not_allowed() { .await .unwrap(); - // Publish with QoS 0 pub_client .publish_qos0("test/upgrade", "Test message") .await @@ -258,7 +237,6 @@ async fn test_qos_upgrade_not_allowed() { sleep(Duration::from_millis(500)).await; - // Message should be received with original QoS 0 (no upgrade) let final_qos = received_qos.load(Ordering::Relaxed); assert_eq!(final_qos, 0, "Message QoS should not be upgraded"); @@ -268,17 +246,14 @@ async fn test_qos_upgrade_not_allowed() { #[tokio::test] async fn test_qos1_retransmission() { - // Start test broker let broker = TestBroker::start().await; - // This test simulates packet loss and retransmission let client = MqttClient::new("qos1-retrans"); client.connect(broker.address()).await.unwrap(); let received = Arc::new(AtomicU32::new(0)); let received_clone = received.clone(); - // Subscribe to our own messages client .subscribe_with_options( "test/retrans", @@ -293,20 +268,19 @@ async fn test_qos1_retransmission() { .await .unwrap(); - // Send a QoS 1 message let result = client .publish_qos1("test/retrans", "Test message") .await .unwrap(); match result { - PublishResult::QoS1Or2 { packet_id } => assert!(packet_id > 0), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result for QoS 1 publish, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => assert!(packet_id > 0), + other => panic!("expected an acknowledged publish, got {other:?}"), } - // Wait for message to arrive sleep(Duration::from_millis(500)).await; - // Should receive exactly one copy assert_eq!(received.load(Ordering::Relaxed), 1); client.disconnect().await.unwrap(); @@ -314,17 +288,14 @@ async fn test_qos1_retransmission() { #[tokio::test] async fn test_qos2_no_duplicates() { - // Start test broker let broker = TestBroker::start().await; - // Test that QoS 2 prevents duplicate delivery let client = MqttClient::new("qos2-nodup"); client.connect(broker.address()).await.unwrap(); let messages = Arc::new(std::sync::Mutex::new(Vec::new())); let messages_clone = messages.clone(); - // Subscribe with QoS 2 client .subscribe_with_options( "test/nodup", @@ -342,7 +313,6 @@ async fn test_qos2_no_duplicates() { .await .unwrap(); - // Send multiple messages with QoS 2 for i in 0..5 { client .publish_qos2("test/nodup", format!("Message {i}")) @@ -350,27 +320,23 @@ async fn test_qos2_no_duplicates() { .unwrap(); } - // Wait for all messages sleep(Duration::from_secs(1)).await; - // Check we received exactly one copy of each { let msgs = messages.lock().unwrap(); assert_eq!(msgs.len(), 5); - // All messages should be unique let mut unique_msgs = msgs.clone(); unique_msgs.sort(); unique_msgs.dedup(); assert_eq!(msgs.len(), unique_msgs.len()); - } // Drop the lock before awaiting + } client.disconnect().await.unwrap(); } #[tokio::test] async fn test_mixed_qos_levels() { - // Start test broker let broker = TestBroker::start().await; let client = MqttClient::new("mixed-qos"); @@ -379,7 +345,6 @@ async fn test_mixed_qos_levels() { let qos_counts = Arc::new(std::sync::Mutex::new([0u32; 3])); let qos_counts_clone = qos_counts.clone(); - // Subscribe with QoS 1 client .subscribe_with_options( "test/mixed", @@ -395,7 +360,6 @@ async fn test_mixed_qos_levels() { .await .unwrap(); - // Send messages with different QoS levels client .publish_qos0("test/mixed", "QoS 0 message") .await @@ -418,7 +382,6 @@ async fn test_mixed_qos_levels() { counts[0], counts[1], counts[2] ); - // Should receive all messages assert_eq!(counts[0], 1, "Should receive QoS 0 message"); assert_eq!( counts[1], 2, @@ -428,20 +391,18 @@ async fn test_mixed_qos_levels() { counts[2], 0, "No messages should be received at QoS 2 (subscription is QoS 1)" ); - } // Drop the lock before awaiting + } client.disconnect().await.unwrap(); } #[tokio::test] async fn test_qos_with_retain() { - // Start test broker let broker = TestBroker::start().await; let pub_client = MqttClient::new("qos-retain-pub"); let sub_client = MqttClient::new("qos-retain-sub"); - // Publisher sends retained message and disconnects pub_client.connect(broker.address()).await.unwrap(); let options = PublishOptions { @@ -456,10 +417,8 @@ async fn test_qos_with_retain() { .unwrap(); pub_client.disconnect().await.unwrap(); - // Wait a bit sleep(Duration::from_millis(100)).await; - // Subscriber connects and subscribes sub_client.connect(broker.address()).await.unwrap(); let received = Arc::new(AtomicU32::new(0)); @@ -481,7 +440,6 @@ async fn test_qos_with_retain() { .await .unwrap(); - // Wait for retained message sleep(Duration::from_millis(500)).await; assert_eq!( @@ -495,24 +453,23 @@ async fn test_qos_with_retain() { #[tokio::test] async fn test_qos_packet_id_exhaustion() { - // Start test broker let broker = TestBroker::start().await; let client = MqttClient::new("qos-exhaustion"); client.connect(broker.address()).await.unwrap(); - // Try to send many QoS 1 messages quickly let mut packet_ids = Vec::new(); - // Send 100 messages rapidly for i in 0..100 { match client .publish_qos1("test/exhaustion", format!("Message {i}")) .await { Ok(result) => match result { - PublishResult::QoS1Or2 { packet_id } => packet_ids.push(packet_id), - PublishResult::QoS0 => panic!("Expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => packet_ids.push(packet_id), + other => panic!("expected an acknowledged publish, got {other:?}"), }, Err(e) => { println!("Failed to send message {i}: {e:?}"); @@ -523,7 +480,6 @@ async fn test_qos_packet_id_exhaustion() { println!("Successfully sent {} messages", packet_ids.len()); - // All successful packet IDs should be unique let mut unique_ids = packet_ids.clone(); unique_ids.sort_unstable(); unique_ids.dedup(); diff --git a/crates/mqtt5/tests/server_max_packet_size.rs b/crates/mqtt5/tests/server_max_packet_size.rs index 3d2c23a4..7440b1f3 100644 --- a/crates/mqtt5/tests/server_max_packet_size.rs +++ b/crates/mqtt5/tests/server_max_packet_size.rs @@ -4,7 +4,7 @@ mod common; use common::{MessageCollector, TestBroker, DEFAULT_TIMEOUT}; use mqtt5::broker::config::{BrokerConfig, StorageBackend, StorageConfig}; use mqtt5::error::MqttError; -use mqtt5::{MqttClient, PublishResult, QoS}; +use mqtt5::{Delivery, MqttClient, PublishResult, QoS}; use std::net::SocketAddr; const BROKER_MAX: usize = 1024; @@ -89,8 +89,10 @@ async fn publish_qos1_packet_id(client: &MqttClient) -> u16 { .await .expect("within-limit QoS1 publish should succeed") { - PublishResult::QoS1Or2 { packet_id } => packet_id, - PublishResult::QoS0 => panic!("expected QoS1Or2 result, got QoS0"), + PublishResult::Sent( + Delivery::AtLeastOnce { packet_id } | Delivery::ExactlyOnce { packet_id }, + ) => packet_id, + other => panic!("expected an acknowledged publish, got {other:?}"), } } diff --git a/crates/mqtt5/tests/shared_subscription_basic.rs b/crates/mqtt5/tests/shared_subscription_basic.rs index 2488ee9c..2b58acdb 100644 --- a/crates/mqtt5/tests/shared_subscription_basic.rs +++ b/crates/mqtt5/tests/shared_subscription_basic.rs @@ -2,9 +2,8 @@ //! Basic test for shared subscriptions use bytes::Bytes; -use mqtt5::broker::router::{DeliveryLanes, LaneReceivers, MessageRouter}; +use mqtt5::broker::router::{DeliveryLanes, LaneReceivers, MessageRouter, SubscriptionRequest}; use mqtt5::packet::publish::PublishPacket; -use mqtt5::types::ProtocolVersion; use mqtt5::QoS; use std::sync::Arc; @@ -20,18 +19,11 @@ async fn register(router: &MessageRouter, client_id: &str, filter: &str) -> Lane ) .await; router - .subscribe( + .subscribe(SubscriptionRequest::new( client_id.to_string(), filter.to_string(), QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); rx diff --git a/crates/mqtt5/tests/turmoil_multi_client.rs b/crates/mqtt5/tests/turmoil_multi_client.rs index 1c364e37..721fcf77 100644 --- a/crates/mqtt5/tests/turmoil_multi_client.rs +++ b/crates/mqtt5/tests/turmoil_multi_client.rs @@ -6,13 +6,12 @@ //! subscription management. #[cfg(feature = "turmoil-testing")] -use mqtt5::broker::router::MessageRouter; +use mqtt5::broker::router::{MessageRouter, SubscriptionRequest}; #[cfg(feature = "turmoil-testing")] use mqtt5::packet::publish::PublishPacket; #[cfg(feature = "turmoil-testing")] use mqtt5::time::Duration; #[cfg(feature = "turmoil-testing")] -use mqtt5::types::ProtocolVersion; #[cfg(feature = "turmoil-testing")] use mqtt5::QoS; #[cfg(feature = "turmoil-testing")] @@ -48,18 +47,11 @@ fn test_multi_client_message_routing() { ) .await; router - .subscribe( - "temp_monitor".to_string(), - "sensors/+/temperature".to_string(), + .subscribe(SubscriptionRequest::new( + "temp_monitor", + "sensors/+/temperature", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -77,18 +69,11 @@ fn test_multi_client_message_routing() { ) .await; router - .subscribe( - "humidity_monitor".to_string(), - "sensors/+/humidity".to_string(), + .subscribe(SubscriptionRequest::new( + "humidity_monitor", + "sensors/+/humidity", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -106,18 +91,11 @@ fn test_multi_client_message_routing() { ) .await; router - .subscribe( - "all_monitor".to_string(), - "sensors/+/+".to_string(), + .subscribe(SubscriptionRequest::new( + "all_monitor", + "sensors/+/+", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -135,18 +113,11 @@ fn test_multi_client_message_routing() { ) .await; router - .subscribe( - "room1_monitor".to_string(), - "sensors/room1/+".to_string(), + .subscribe(SubscriptionRequest::new( + "room1_monitor", + "sensors/room1/+", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -240,18 +211,11 @@ fn test_client_subscription_changes() { // Initial subscription router - .subscribe( - "dynamic_client".to_string(), - "alerts/error".to_string(), + .subscribe(SubscriptionRequest::new( + "dynamic_client", + "alerts/error", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -282,18 +246,11 @@ fn test_client_subscription_changes() { // Add another subscription for warnings router - .subscribe( - "dynamic_client".to_string(), - "alerts/warning".to_string(), + .subscribe(SubscriptionRequest::new( + "dynamic_client", + "alerts/warning", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -350,34 +307,20 @@ fn test_message_ordering_with_multiple_clients() { .await; router - .subscribe( - "client1".to_string(), - "sequence/test".to_string(), + .subscribe(SubscriptionRequest::new( + "client1", + "sequence/test", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "client2".to_string(), - "sequence/test".to_string(), + .subscribe(SubscriptionRequest::new( + "client2", + "sequence/test", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); diff --git a/crates/mqtt5/tests/turmoil_pubsub.rs b/crates/mqtt5/tests/turmoil_pubsub.rs index 2771d835..9dff5a61 100644 --- a/crates/mqtt5/tests/turmoil_pubsub.rs +++ b/crates/mqtt5/tests/turmoil_pubsub.rs @@ -5,13 +5,12 @@ //! environment, testing various `QoS` levels, topic patterns, and edge cases. #[cfg(feature = "turmoil-testing")] -use mqtt5::broker::router::MessageRouter; +use mqtt5::broker::router::{MessageRouter, SubscriptionRequest}; #[cfg(feature = "turmoil-testing")] use mqtt5::packet::publish::PublishPacket; #[cfg(feature = "turmoil-testing")] use mqtt5::time::Duration; #[cfg(feature = "turmoil-testing")] -use mqtt5::types::ProtocolVersion; #[cfg(feature = "turmoil-testing")] use mqtt5::QoS; #[cfg(feature = "turmoil-testing")] @@ -39,18 +38,11 @@ fn test_basic_publish_subscribe() { ) .await; router - .subscribe( - "subscriber".to_string(), - "test/topic".to_string(), + .subscribe(SubscriptionRequest::new( + "subscriber", + "test/topic", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -115,35 +107,21 @@ fn test_wildcard_subscriptions() { // Single-level wildcard subscription router - .subscribe( - "single_wildcard".to_string(), - "sensors/+/temperature".to_string(), + .subscribe(SubscriptionRequest::new( + "single_wildcard", + "sensors/+/temperature", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); // Multi-level wildcard subscription router - .subscribe( - "multi_wildcard".to_string(), - "sensors/#".to_string(), + .subscribe(SubscriptionRequest::new( + "multi_wildcard", + "sensors/#", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -248,18 +226,11 @@ fn test_multiple_subscribers_same_topic() { let topic = "broadcast/announcement"; for client in ["subscriber1", "subscriber2", "subscriber3"] { router - .subscribe( + .subscribe(SubscriptionRequest::new( client.to_string(), topic.to_string(), QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); } @@ -336,34 +307,20 @@ fn test_qos_levels() { // Subscribe with different QoS levels router - .subscribe( - "qos0_client".to_string(), - "data/qos0".to_string(), + .subscribe(SubscriptionRequest::new( + "qos0_client", + "data/qos0", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); router - .subscribe( - "qos1_client".to_string(), - "data/qos1".to_string(), + .subscribe(SubscriptionRequest::new( + "qos1_client", + "data/qos1", QoS::AtLeastOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); @@ -427,18 +384,11 @@ fn test_unsubscribe_functionality() { // Subscribe to topic router - .subscribe( - "test_client".to_string(), - "test/unsubscribe".to_string(), + .subscribe(SubscriptionRequest::new( + "test_client", + "test/unsubscribe", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); diff --git a/crates/mqtt5/tests/turmoil_shared_subscriptions.rs b/crates/mqtt5/tests/turmoil_shared_subscriptions.rs index 2fc1b51c..7838eaaa 100644 --- a/crates/mqtt5/tests/turmoil_shared_subscriptions.rs +++ b/crates/mqtt5/tests/turmoil_shared_subscriptions.rs @@ -4,10 +4,9 @@ //! These tests verify the shared subscription functionality using //! the existing `MessageRouter` directly, which we know works. -use mqtt5::broker::router::{DeliveryLanes, LaneReceivers, MessageRouter}; +use mqtt5::broker::router::{DeliveryLanes, LaneReceivers, MessageRouter, SubscriptionRequest}; use mqtt5::packet::publish::PublishPacket; use mqtt5::time::Duration; -use mqtt5::types::ProtocolVersion; use mqtt5::QoS; use std::sync::Arc; @@ -23,18 +22,11 @@ async fn register_worker(router: &MessageRouter, client_id: &str) -> LaneReceive ) .await; router - .subscribe( + .subscribe(SubscriptionRequest::new( client_id.to_string(), - "$share/workers/tasks/+".to_string(), + "$share/workers/tasks/+", QoS::AtMostOnce, - None, - false, - false, - 0, - ProtocolVersion::V5, - false, - None, - ) + )) .await .unwrap(); rx diff --git a/crates/mqttv5-cli/src/commands/bench_cmd.rs b/crates/mqttv5-cli/src/commands/bench_cmd.rs index e816a221..cf10e79b 100644 --- a/crates/mqttv5-cli/src/commands/bench_cmd.rs +++ b/crates/mqttv5-cli/src/commands/bench_cmd.rs @@ -1003,12 +1003,12 @@ async fn publish_message( payload: Vec, qos: QoS, ) -> Result<()> { - match qos { + let result = match qos { QoS::AtMostOnce => client.publish(topic, payload).await?, QoS::AtLeastOnce => client.publish_qos1(topic, payload).await?, QoS::ExactlyOnce => client.publish_qos2(topic, payload).await?, }; - Ok(()) + super::pub_cmd::require_sent(&result) } struct PayloadSpec { diff --git a/crates/mqttv5-cli/src/commands/pub_cmd.rs b/crates/mqttv5-cli/src/commands/pub_cmd.rs index a1a987ec..56d66df9 100644 --- a/crates/mqttv5-cli/src/commands/pub_cmd.rs +++ b/crates/mqttv5-cli/src/commands/pub_cmd.rs @@ -6,7 +6,8 @@ use mqtt5::time::Duration; #[cfg(feature = "codec")] use mqtt5::{CodecRegistry, DeflateCodec, GzipCodec}; use mqtt5::{ - ConnectOptions, ConnectionEvent, Message, MqttClient, PublishOptions, QoS, WillMessage, + ConnectOptions, ConnectionEvent, Message, MqttClient, PublishOptions, PublishResult, QoS, + WillMessage, }; use std::io::{self, Read}; use std::sync::atomic::{AtomicU32, Ordering}; @@ -567,25 +568,29 @@ async fn publish_with_properties( .response_topic .clone_from(&cmd.response.response_topic); options.properties.correlation_data = correlation_data.cloned(); - client - .publish_with_options(topic, message.as_bytes(), options) - .await?; - Ok(()) + require_sent( + &client + .publish_with_options(topic, message.as_bytes(), options) + .await?, + ) } async fn publish_simple(client: &MqttClient, topic: &str, message: &str, qos: QoS) -> Result<()> { - match qos { - QoS::AtMostOnce => { - client.publish(topic, message.as_bytes()).await?; - } - QoS::AtLeastOnce => { - client.publish_qos1(topic, message.as_bytes()).await?; - } - QoS::ExactlyOnce => { - client.publish_qos2(topic, message.as_bytes()).await?; - } + let result = match qos { + QoS::AtMostOnce => client.publish(topic, message.as_bytes()).await?, + QoS::AtLeastOnce => client.publish_qos1(topic, message.as_bytes()).await?, + QoS::ExactlyOnce => client.publish_qos2(topic, message.as_bytes()).await?, + }; + require_sent(&result) +} + +pub(crate) fn require_sent(result: &PublishResult) -> Result<()> { + match result { + PublishResult::Sent(_) => Ok(()), + PublishResult::Queued(_) => Err(anyhow::anyhow!( + "publish was not acknowledged before the connection ended or the acknowledgement wait elapsed" + )), } - Ok(()) } fn print_publish_result(cmd: &PubCommand, iteration: u64, topic: &str, qos: QoS) { diff --git a/specs/tla/offline-queue/OfflineQueue.cfg b/specs/tla/offline-queue/OfflineQueue.cfg new file mode 100644 index 00000000..0e5aac51 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue.cfg @@ -0,0 +1,37 @@ +SPECIFICATION Spec + +CONSTANTS + N = 3 + K = 2 + MaxConns = 3 + RMSet = {1} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvRetainAvailable + InvMaximumQoS + InvMaximumPacketSize + InvReceiveMaximum + InvNoSilentLoss + InvRejectedNeverDelivered + InvOrder + InvPidUnique + InvPidNotQuarantined + InvNoStaleServerPid + InvNoMisattribution + InvExactlyOnce + InvRetainFidelity + InvQoSFidelity + InvFairActionEnabled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue.tla b/specs/tla/offline-queue/OfflineQueue.tla new file mode 100644 index 00000000..5f9da0f1 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue.tla @@ -0,0 +1,523 @@ +---------------------------- MODULE OfflineQueue ---------------------------- +EXTENDS Naturals, Sequences, FiniteSets + +CONSTANTS + N, + K, + MaxConns, + RMSet, + QSet, + RSet, + BSet, + MQSet, + RASet, + MBSet, + Design, + Quarantine, + Report, + ReplayDowngrade + +Msgs == 1..N +Ids == 1..K +Epochs == 1..MaxConns +Outcomes == {"none", "ok", "rejected", "indet"} +Pcs == {"none", "replay", "flush", "conformed", "popped", "stored", "done"} + +Fix == Design = "Fix" + +MaxOf(S) == CHOOSE x \in S : \A y \in S : x >= y +MinOf(S) == CHOOSE x \in S : \A y \in S : x <= y +Min(a, b) == IF a < b THEN a ELSE b + +Full == [mq |-> 2, ra |-> TRUE, mb |-> TRUE, rm |-> MaxOf(RMSet)] +NoCur == [m |-> 0, id |-> 0, q |-> 0, re |-> FALSE] +IdleTask == [pc |-> "none", snap |-> <<>>, pol |-> Full, cur |-> NoCur, bad |-> FALSE] + +KOk == [x \in Msgs |-> "ok"] +KRejected == [x \in Msgs |-> "rejected"] +KIndet == [x \in Msgs |-> "indet"] + +VARIABLES + reqQ, reqR, big, + nextM, acc, outcome, outQ, effQ, + callerMap, events, misattr, + queue, store, task, quar, infl, + up, conns, caps, + wire, acks, + rcount, rlog, srvQ2 + +reportVars == <> +vars == <> + +Range(s) == {s[i] : i \in DOMAIN s} +MsgSet(s) == {s[i].m : i \in DOMAIN s} +RemoveMsg(s, m) == SelectSeq(s, LAMBDA x : x.m # m) +StoreDel(s, i) == SelectSeq(s, LAMBDA r : r.id # i) +StoreSet(s, i, f, v) == + [j \in DOMAIN s |-> IF s[j].id = i THEN [s[j] EXCEPT ![f] = v] ELSE s[j]] + +CurTasks == {e \in Epochs : task[e].pc \in {"conformed", "popped", "stored"}} +PoppedTasks == {e \in Epochs : task[e].pc \in {"popped", "stored"}} +HeldCurs == {task[e].cur : e \in {x \in CurTasks : task[x].cur.q > 0}} + +Holders == + {<> : x \in Range(queue)} + \cup {<> : r \in Range(store)} + \cup {<> : c \in HeldCurs} + +InUse == + {x.id : x \in Range(queue)} + \cup {r.id : r \in Range(store)} + \cup (IF Fix THEN {c.id : c \in HeldCurs} \cup quar ELSE {}) + +IsLive(m) == + \/ \E x \in Range(queue) : x.m = m + \/ \E r \in Range(store) : r.m = m + \/ \E e \in PoppedTasks : task[e].cur.m = m + +IsPub(i) == wire[i].k = "pub" +OnWire(m) == \E i \in DOMAIN wire : wire[i].m = m +PendingEvent(m) == \E i \in DOMAIN events : events[i].m = m + +QFor(x, mode) == + CASE mode = "eff" -> effQ[x] + [] mode = "req" -> reqQ[x] + [] OTHER -> 0 + +Resolve(s, kf, mode) == + IF Report = "handle" + THEN /\ outcome' = [x \in Msgs |-> + IF x \in MsgSet(s) /\ outcome[x] = "none" THEN kf[x] ELSE outcome[x]] + /\ outQ' = [x \in Msgs |-> + IF x \in MsgSet(s) /\ outcome[x] = "none" THEN QFor(x, mode) ELSE outQ[x]] + /\ UNCHANGED events + ELSE /\ events' = events \o [i \in 1..Len(s) |-> + [id |-> s[i].id, k |-> kf[s[i].m], + q |-> QFor(s[i].m, mode), m |-> s[i].m]] + /\ UNCHANGED <> + +PubPacket(m, id, q) == [k |-> "pub", m |-> m, id |-> IF q = 0 THEN 0 ELSE id, q |-> q, r |-> reqR[m]] +RelPacket(m, id) == [k |-> "rel", m |-> m, id |-> id, q |-> 2, r |-> FALSE] + +Init == + /\ reqQ \in [Msgs -> QSet] + /\ reqR \in [Msgs -> RSet] + /\ big \in [Msgs -> BSet] + /\ nextM = 1 + /\ acc = [m \in Msgs |-> FALSE] + /\ outcome = [m \in Msgs |-> "none"] + /\ outQ = [m \in Msgs |-> 0] + /\ effQ = reqQ + /\ callerMap = [i \in Ids |-> 0] + /\ events = <<>> + /\ misattr = FALSE + /\ queue = <<>> + /\ store = <<>> + /\ task = [e \in Epochs |-> IdleTask] + /\ quar = {} + /\ infl = {} + /\ up = FALSE + /\ conns = 0 + /\ caps = Full + /\ wire = <<>> + /\ acks = <<>> + /\ rcount = [m \in Msgs |-> 0] + /\ rlog = <<>> + /\ srvQ2 = {} + +Accept == + /\ nextM <= N + /\ LET m == nextM + violatesLastKnown == + \/ reqR[m] /\ ~caps.ra /\ (Fix \/ up) + \/ big[m] /\ ~caps.mb + free == Ids \ InUse + IN /\ nextM' = nextM + 1 + /\ IF violatesLastKnown \/ free = {} + THEN /\ outcome' = [outcome EXCEPT ![m] = "rejected"] + /\ UNCHANGED <> + ELSE LET p == MinOf(free) IN + /\ acc' = [acc EXCEPT ![m] = TRUE] + /\ queue' = Append(queue, [m |-> m, id |-> p, re |-> FALSE]) + /\ callerMap' = IF Report = "pidEvent" + THEN [callerMap EXCEPT ![p] = m] + ELSE callerMap + /\ IF Design = "D0" /\ ~up + THEN /\ outcome' = [outcome EXCEPT ![m] = "ok"] + /\ outQ' = [outQ EXCEPT ![m] = reqQ[m]] + ELSE UNCHANGED <> + /\ task' = IF up /\ task[conns].pc = "done" + THEN [task EXCEPT ![conns].pc = "flush"] + ELSE task + /\ UNCHANGED <> + +Connect == + /\ ~up + /\ conns < MaxConns + /\ \E mq \in MQSet, ra \in RASet, mb \in MBSet, rm \in RMSet, + sp \in (IF conns = 0 THEN {FALSE} ELSE BOOLEAN) : + LET c == [mq |-> mq, ra |-> ra, mb |-> mb, rm |-> rm] + e == conns + 1 + fresh == [pc |-> "replay", + snap |-> IF sp THEN [i \in 1..Len(store) |-> store[i].id] ELSE <<>>, + pol |-> c, cur |-> NoCur, bad |-> FALSE] + qos1 == SelectSeq(store, LAMBDA r : r.q = 1) + requeue == [i \in 1..Len(qos1) |-> [m |-> qos1[i].m, id |-> qos1[i].id, re |-> TRUE]] + qos2 == SelectSeq(store, LAMBDA r : r.q = 2) + owned == {r.m : r \in {x \in Range(store) : x.rel}} + qos2Kind == [x \in Msgs |-> IF x \in owned THEN "ok" ELSE "indet"] + IN /\ caps' = c + /\ conns' = e + /\ task' = [x \in Epochs |-> + CASE x = e -> fresh + [] Fix -> IdleTask + [] OTHER -> task[x]] + /\ IF sp + THEN UNCHANGED <> + ELSE /\ srvQ2' = {} + /\ quar' = {} + /\ store' = <<>> + /\ IF Fix + THEN /\ queue' = requeue \o queue + /\ Resolve(qos2, qos2Kind, "eff") + ELSE /\ UNCHANGED queue + /\ Resolve(store, KIndet, "eff") + /\ up' = TRUE + /\ infl' = {} + /\ UNCHANGED <> + +Lose == + /\ up + /\ conns < MaxConns + /\ up' = FALSE + /\ wire' = <<>> + /\ acks' = <<>> + /\ infl' = {} + /\ UNCHANGED <> + +TaskFrame == <> + +SendReplay(e, r, q) == + /\ q > 0 => Cardinality(infl) < caps.rm + /\ task' = [task EXCEPT ![e].snap = Tail(@)] + /\ wire' = Append(wire, PubPacket(r.m, r.id, q)) + /\ effQ' = [effQ EXCEPT ![r.m] = Min(@, q)] + /\ IF q = 0 + THEN /\ store' = StoreDel(store, r.id) + /\ Resolve(<>, KOk, "zero") + /\ UNCHANGED infl + ELSE /\ store' = StoreSet(store, r.id, "q", q) + /\ infl' = infl \cup {r.id} + /\ UNCHANGED reportVars + /\ UNCHANGED quar + +SendReplayRel(e, r) == + /\ Cardinality(infl) < caps.rm + /\ task' = [task EXCEPT ![e].snap = Tail(@)] + /\ wire' = Append(wire, RelPacket(r.m, r.id)) + /\ infl' = infl \cup {r.id} + /\ UNCHANGED <> + +AbandonReplay(e, r) == + /\ task' = [task EXCEPT ![e].snap = Tail(@)] + /\ store' = StoreDel(store, r.id) + /\ quar' = IF Quarantine /\ r.q = 2 THEN quar \cup {r.id} ELSE quar + /\ Resolve(<>, KIndet, "eff") + /\ UNCHANGED <> + +ReplayStep(e) == + /\ task[e].pc = "replay" + /\ task[e].snap # <<>> + /\ e = conns + /\ up + /\ LET i == Head(task[e].snap) + hits == {r \in Range(store) : r.id = i} + IN IF hits = {} + THEN /\ task' = [task EXCEPT ![e].snap = Tail(@)] + /\ UNCHANGED <> + ELSE LET r == CHOOSE x \in hits : TRUE + shapeOk == ~(reqR[r.m] /\ ~caps.ra) /\ ~(big[r.m] /\ ~caps.mb) + IN CASE r.rel -> SendReplayRel(e, r) + [] Design = "D0" \/ (shapeOk /\ r.q <= caps.mq) -> SendReplay(e, r, r.q) + [] ReplayDowngrade /\ shapeOk -> SendReplay(e, r, caps.mq) + [] OTHER -> AbandonReplay(e, r) + /\ UNCHANGED queue + /\ UNCHANGED TaskFrame + +ReplayEnd(e) == + /\ task[e].pc = "replay" + /\ task[e].snap = <<>> + /\ task' = [task EXCEPT ![e].pc = "flush"] + /\ UNCHANGED <> + /\ UNCHANGED TaskFrame + +ReplayDie(e) == + /\ task[e].pc = "replay" + /\ task[e].snap # <<>> + /\ e # conns \/ ~up + /\ task' = [task EXCEPT ![e] = IdleTask] + /\ UNCHANGED <> + /\ UNCHANGED TaskFrame + +FlushFinish(e) == + /\ task[e].pc = "flush" + /\ queue = <<>> + /\ task' = [task EXCEPT ![e].pc = "done"] + /\ UNCHANGED <> + /\ UNCHANGED TaskFrame + +Conform(e) == + /\ task[e].pc = "flush" + /\ queue # <<>> + /\ Fix => (e = conns /\ up) + /\ LET h == Head(queue) + pol == task[e].pol + IN task' = [task EXCEPT ![e].pc = "conformed", + ![e].cur = [m |-> h.m, id |-> h.id, + q |-> Min(effQ[h.m], pol.mq), re |-> h.re], + ![e].bad = (reqR[h.m] /\ ~pol.ra) \/ (big[h.m] /\ ~pol.mb)] + /\ UNCHANGED <> + /\ UNCHANGED TaskFrame + +Drop(e) == + /\ task[e].pc = "conformed" + /\ task[e].bad + /\ Fix => e = conns + /\ LET c == task[e].cur IN + IF Fix + THEN /\ queue' = RemoveMsg(queue, c.m) + /\ Resolve(<>, IF c.re THEN KIndet ELSE KRejected, "eff") + ELSE /\ queue' = IF queue = <<>> THEN queue ELSE Tail(queue) + /\ UNCHANGED reportVars + /\ task' = [task EXCEPT ![e].pc = "flush", ![e].cur = NoCur, ![e].bad = FALSE] + /\ UNCHANGED <> + /\ UNCHANGED TaskFrame + +ClaimFix(e) == + /\ e = conns + /\ up + /\ LET p == task[e].cur IN + /\ p.q > 0 => Cardinality(infl) < caps.rm + /\ queue' = RemoveMsg(queue, p.m) + /\ wire' = Append(wire, PubPacket(p.m, p.id, p.q)) + /\ effQ' = [effQ EXCEPT ![p.m] = Min(@, p.q)] + /\ task' = [task EXCEPT ![e].pc = "flush", ![e].cur = NoCur] + /\ IF p.q = 0 + THEN /\ Resolve(<

>, KOk, "zero") + /\ UNCHANGED <> + ELSE /\ store' = Append(store, [m |-> p.m, id |-> p.id, q |-> p.q, rel |-> FALSE]) + /\ infl' = infl \cup {p.id} + /\ UNCHANGED reportVars + +ClaimD0(e) == + /\ LET p == task[e].cur IN + /\ p.q > 0 => (e = conns /\ Cardinality(infl) < caps.rm) + /\ queue' = IF queue = <<>> THEN queue ELSE Tail(queue) + /\ infl' = IF p.q = 0 THEN infl ELSE infl \cup {p.id} + /\ task' = [task EXCEPT ![e].pc = "popped"] + /\ UNCHANGED <> + +Claim(e) == + /\ task[e].pc = "conformed" + /\ ~task[e].bad + /\ IF Fix THEN ClaimFix(e) ELSE ClaimD0(e) + /\ UNCHANGED quar + /\ UNCHANGED TaskFrame + +StoreStep(e) == + /\ task[e].pc = "popped" + /\ LET p == task[e].cur IN + store' = IF p.q > 0 + THEN Append(StoreDel(store, p.id), [m |-> p.m, id |-> p.id, q |-> p.q, rel |-> FALSE]) + ELSE store + /\ task' = [task EXCEPT ![e].pc = "stored"] + /\ UNCHANGED <> + /\ UNCHANGED TaskFrame + +WriteStep(e) == + /\ task[e].pc = "stored" + /\ LET p == task[e].cur IN + IF e = conns /\ up + THEN /\ wire' = Append(wire, PubPacket(p.m, p.id, p.q)) + /\ effQ' = [effQ EXCEPT ![p.m] = Min(@, p.q)] + /\ task' = [task EXCEPT ![e].pc = "flush", ![e].cur = NoCur] + /\ IF p.q = 0 THEN Resolve(<

>, KOk, "zero") ELSE UNCHANGED reportVars + ELSE /\ task' = [task EXCEPT ![e] = IdleTask] + /\ UNCHANGED <> + /\ UNCHANGED <> + /\ UNCHANGED TaskFrame + +TaskStep(e) == + \/ ReplayStep(e) \/ ReplayEnd(e) \/ ReplayDie(e) + \/ FlushFinish(e) \/ Conform(e) \/ Drop(e) \/ Claim(e) + \/ StoreStep(e) \/ WriteStep(e) + +TaskAny == \E e \in Epochs : TaskStep(e) + +ServerRecv == + /\ up + /\ wire # <<>> + /\ LET p == Head(wire) + dedup == p.q = 2 /\ \E pr \in srvQ2 : pr[1] = p.id + IN /\ wire' = Tail(wire) + /\ IF p.k = "rel" + THEN /\ srvQ2' = {pr \in srvQ2 : pr[1] # p.id} + /\ acks' = Append(acks, [k |-> "comp", id |-> p.id]) + /\ UNCHANGED <> + ELSE /\ rcount' = IF dedup THEN rcount ELSE [rcount EXCEPT ![p.m] = Min(@ + 1, 2)] + /\ rlog' = IF ~dedup /\ rcount[p.m] = 0 THEN Append(rlog, p.m) ELSE rlog + /\ srvQ2' = IF p.q = 2 /\ ~dedup THEN srvQ2 \cup {<>} ELSE srvQ2 + /\ acks' = CASE p.q = 2 -> Append(acks, [k |-> "rec", id |-> p.id]) + [] p.q = 1 -> Append(acks, [k |-> "ack", id |-> p.id]) + [] OTHER -> acks + /\ UNCHANGED <> + +ClientAck == + /\ up + /\ acks # <<>> + /\ LET a == Head(acks) + hits == {r \in Range(store) : r.id = a.id} + IN /\ acks' = Tail(acks) + /\ IF hits = {} + THEN UNCHANGED <> + ELSE LET r == CHOOSE x \in hits : TRUE IN + IF a.k = "rec" + THEN /\ store' = StoreSet(store, a.id, "rel", TRUE) + /\ wire' = Append(wire, RelPacket(r.m, a.id)) + /\ UNCHANGED <> + ELSE /\ store' = StoreDel(store, a.id) + /\ infl' = infl \ {a.id} + /\ Resolve(<>, KOk, IF Fix THEN "eff" ELSE "req") + /\ UNCHANGED wire + /\ UNCHANGED <> + +CallerConsume == + /\ events # <<>> + /\ LET ev == Head(events) + t == callerMap[ev.id] + IN /\ events' = Tail(events) + /\ outcome' = IF outcome[t] = "none" THEN [outcome EXCEPT ![t] = ev.k] ELSE outcome + /\ outQ' = IF outcome[t] = "none" THEN [outQ EXCEPT ![t] = ev.q] ELSE outQ + /\ misattr' = (misattr \/ t # ev.m) + /\ UNCHANGED <> + +Next == + \/ Accept + \/ Connect + \/ Lose + \/ TaskAny + \/ ServerRecv + \/ ClientAck + \/ CallerConsume + +Spec == + /\ Init + /\ [][Next]_vars + /\ WF_vars(Connect) + /\ WF_vars(TaskAny) + /\ WF_vars(ServerRecv) + /\ WF_vars(ClientAck) + /\ WF_vars(CallerConsume) + +SpecNoConnectFairness == + /\ Init + /\ [][Next]_vars + /\ WF_vars(TaskAny) + /\ WF_vars(ServerRecv) + /\ WF_vars(ClientAck) + /\ WF_vars(CallerConsume) + +---------------------------------------------------------------------------- + +TypeOK == + /\ reqQ \in [Msgs -> 1..2] + /\ nextM \in 1..(N + 1) + /\ acc \in [Msgs -> BOOLEAN] + /\ outcome \in [Msgs -> Outcomes] + /\ outQ \in [Msgs -> 0..2] + /\ effQ \in [Msgs -> 0..2] + /\ callerMap \in [Ids -> 0..N] + /\ Len(events) <= N + /\ misattr \in BOOLEAN + /\ Len(queue) <= N + /\ Len(store) <= N + /\ \A e \in Epochs : task[e].pc \in Pcs + /\ quar \subseteq Ids + /\ infl \subseteq Ids + /\ up \in BOOLEAN + /\ conns \in 0..MaxConns + /\ Len(wire) <= 3 * N + /\ rcount \in [Msgs -> 0..2] + /\ Len(rlog) <= N + +InvRetainAvailable == \A i \in DOMAIN wire : (IsPub(i) /\ wire[i].r) => caps.ra + +InvMaximumQoS == \A i \in DOMAIN wire : IsPub(i) => wire[i].q <= caps.mq + +InvMaximumPacketSize == \A i \in DOMAIN wire : (IsPub(i) /\ big[wire[i].m]) => caps.mb + +InvReceiveMaximum == Cardinality(infl) <= caps.rm + +InvNoSilentLoss == + \A m \in Msgs : + (acc[m] /\ ~IsLive(m) /\ rcount[m] = 0 /\ ~OnWire(m)) => + \/ outcome[m] \in {"rejected", "indet"} + \/ outcome[m] = "ok" /\ outQ[m] = 0 + \/ PendingEvent(m) + +InvRejectedNeverDelivered == + \A m \in Msgs : outcome[m] = "rejected" => (rcount[m] = 0 /\ ~OnWire(m)) + +InvOrder == \A i, j \in DOMAIN rlog : i < j => rlog[i] < rlog[j] + +InvPidUnique == \A a, b \in Holders : a[2] = b[2] => a[1] = b[1] + +InvPidNotQuarantined == \A h \in Holders : h[2] \notin quar + +InvNoStaleServerPid == + /\ \A i \in DOMAIN wire : \A pr \in srvQ2 : + pr[1] = wire[i].id => (pr[2] = wire[i].m /\ wire[i].q = 2) + /\ \A h \in Holders : \A pr \in srvQ2 : pr[1] = h[2] => pr[2] = h[1] + +InvNoMisattribution == ~misattr + +InvExactlyOnce == + \A m \in Msgs : + rcount[m] > 1 => + /\ effQ[m] < 2 + /\ outcome[m] \in {"ok", "indet"} => outQ[m] < 2 + +InvRetainFidelity == \A i \in DOMAIN wire : IsPub(i) => wire[i].r = reqR[wire[i].m] + +InvQoSFidelity == \A m \in Msgs : outcome[m] \in {"ok", "indet"} => outQ[m] = effQ[m] + +Unresolved == \E m \in Msgs : acc[m] /\ outcome[m] = "none" + +InvFairActionEnabled == + Unresolved => ENABLED (Connect \/ TaskAny \/ ServerRecv \/ ClientAck \/ CallerConsume) + +AllResolved == \A m \in Msgs : acc[m] => outcome[m] # "none" + +Settled == queue = <<>> /\ store = <<>> + +LiveResolved == []<>AllResolved + +LiveSettled == []<>Settled + +NegAllDelivered == []<>(\A m \in Msgs : acc[m] => outcome[m] = "ok") + +NegQuarantineReleased == []<>(quar = {}) + +NegImpossible == []<>(nextM > N + 1) +============================================================================= diff --git a/specs/tla/offline-queue/OfflineQueue_D0_ExactlyOnce.cfg b/specs/tla/offline-queue/OfflineQueue_D0_ExactlyOnce.cfg new file mode 100644 index 00000000..00b2f65c --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_ExactlyOnce.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {2} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {1, 2} + RASet = {TRUE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvExactlyOnce + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_MaximumPacketSize.cfg b/specs/tla/offline-queue/OfflineQueue_D0_MaximumPacketSize.cfg new file mode 100644 index 00000000..9a18c844 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_MaximumPacketSize.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {1} + RSet = {FALSE} + BSet = {TRUE} + MQSet = {1} + RASet = {TRUE} + MBSet = {TRUE, FALSE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvMaximumPacketSize + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_MaximumQoS.cfg b/specs/tla/offline-queue/OfflineQueue_D0_MaximumQoS.cfg new file mode 100644 index 00000000..8993e2a7 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_MaximumQoS.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {2} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {1, 2} + RASet = {TRUE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvMaximumQoS + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_NoSilentLoss.cfg b/specs/tla/offline-queue/OfflineQueue_D0_NoSilentLoss.cfg new file mode 100644 index 00000000..864f9417 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_NoSilentLoss.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {1} + RSet = {FALSE} + BSet = {TRUE, FALSE} + MQSet = {1} + RASet = {TRUE} + MBSet = {TRUE, FALSE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvNoSilentLoss + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_NoStaleServerPid.cfg b/specs/tla/offline-queue/OfflineQueue_D0_NoStaleServerPid.cfg new file mode 100644 index 00000000..cdd3c5c5 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_NoStaleServerPid.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {2} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {2} + RASet = {TRUE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvNoStaleServerPid + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_Order.cfg b/specs/tla/offline-queue/OfflineQueue_D0_Order.cfg new file mode 100644 index 00000000..4e3f3441 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_Order.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 3 + RMSet = {1} + QSet = {1} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {1} + RASet = {TRUE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvOrder + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_PidUnique.cfg b/specs/tla/offline-queue/OfflineQueue_D0_PidUnique.cfg new file mode 100644 index 00000000..5ffb6dff --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_PidUnique.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {1} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {1} + RASet = {TRUE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvPidUnique + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_QoSFidelity.cfg b/specs/tla/offline-queue/OfflineQueue_D0_QoSFidelity.cfg new file mode 100644 index 00000000..9bfd9af9 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_QoSFidelity.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {2} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {1, 2} + RASet = {TRUE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvQoSFidelity + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_RetainAvailable.cfg b/specs/tla/offline-queue/OfflineQueue_D0_RetainAvailable.cfg new file mode 100644 index 00000000..cbc93874 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_RetainAvailable.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {1} + RSet = {TRUE} + BSet = {FALSE} + MQSet = {1} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvRetainAvailable + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_D0_live.cfg b/specs/tla/offline-queue/OfflineQueue_D0_live.cfg new file mode 100644 index 00000000..637ac226 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_D0_live.cfg @@ -0,0 +1,25 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 3 + RMSet = {1} + QSet = {1} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {1} + RASet = {TRUE} + MBSet = {TRUE} + Design = "D0" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + +PROPERTIES + LiveSettled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_NEG_alldelivered.cfg b/specs/tla/offline-queue/OfflineQueue_NEG_alldelivered.cfg new file mode 100644 index 00000000..65ed85a8 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_NEG_alldelivered.cfg @@ -0,0 +1,25 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 3 + RMSet = {1, 2} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {0, 2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + +PROPERTIES + NegAllDelivered + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_NEG_impossible.cfg b/specs/tla/offline-queue/OfflineQueue_NEG_impossible.cfg new file mode 100644 index 00000000..660bfb4b --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_NEG_impossible.cfg @@ -0,0 +1,25 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 3 + RMSet = {1, 2} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {0, 2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + +PROPERTIES + NegImpossible + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_NEG_noconnectfair.cfg b/specs/tla/offline-queue/OfflineQueue_NEG_noconnectfair.cfg new file mode 100644 index 00000000..72ebaa1c --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_NEG_noconnectfair.cfg @@ -0,0 +1,25 @@ +SPECIFICATION SpecNoConnectFairness + +CONSTANTS + N = 2 + K = 2 + MaxConns = 3 + RMSet = {1, 2} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {0, 2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + +PROPERTIES + LiveResolved + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_NEG_quarantine.cfg b/specs/tla/offline-queue/OfflineQueue_NEG_quarantine.cfg new file mode 100644 index 00000000..862b8a44 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_NEG_quarantine.cfg @@ -0,0 +1,25 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 3 + RMSet = {1, 2} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {0, 2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + +PROPERTIES + NegQuarantineReleased + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_V_noquarantine.cfg b/specs/tla/offline-queue/OfflineQueue_V_noquarantine.cfg new file mode 100644 index 00000000..3c4ee740 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_V_noquarantine.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = FALSE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvNoStaleServerPid + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_V_pidevent.cfg b/specs/tla/offline-queue/OfflineQueue_V_pidevent.cfg new file mode 100644 index 00000000..499382c5 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_V_pidevent.cfg @@ -0,0 +1,23 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {1} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {1} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "pidEvent" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvNoMisattribution + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_V_replaydowngrade.cfg b/specs/tla/offline-queue/OfflineQueue_V_replaydowngrade.cfg new file mode 100644 index 00000000..d6d81f63 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_V_replaydowngrade.cfg @@ -0,0 +1,37 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1} + QSet = {2} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {1, 2} + RASet = {TRUE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = TRUE + +INVARIANTS + TypeOK + InvRetainAvailable + InvMaximumQoS + InvMaximumPacketSize + InvReceiveMaximum + InvNoSilentLoss + InvRejectedNeverDelivered + InvOrder + InvPidUnique + InvPidNotQuarantined + InvNoStaleServerPid + InvNoMisattribution + InvExactlyOnce + InvRetainFidelity + InvQoSFidelity + InvFairActionEnabled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_alldims.cfg b/specs/tla/offline-queue/OfflineQueue_alldims.cfg new file mode 100644 index 00000000..bc7296d8 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_alldims.cfg @@ -0,0 +1,37 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 2 + RMSet = {1, 2} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {TRUE, FALSE} + MQSet = {0, 1, 2} + RASet = {TRUE, FALSE} + MBSet = {TRUE, FALSE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvRetainAvailable + InvMaximumQoS + InvMaximumPacketSize + InvReceiveMaximum + InvNoSilentLoss + InvRejectedNeverDelivered + InvOrder + InvPidUnique + InvPidNotQuarantined + InvNoStaleServerPid + InvNoMisattribution + InvExactlyOnce + InvRetainFidelity + InvQoSFidelity + InvFairActionEnabled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_live.cfg b/specs/tla/offline-queue/OfflineQueue_live.cfg new file mode 100644 index 00000000..b0c03fcc --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_live.cfg @@ -0,0 +1,26 @@ +SPECIFICATION Spec + +CONSTANTS + N = 2 + K = 2 + MaxConns = 3 + RMSet = {1, 2} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {0, 2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + +PROPERTIES + LiveResolved + LiveSettled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_qos.cfg b/specs/tla/offline-queue/OfflineQueue_qos.cfg new file mode 100644 index 00000000..fcb33be5 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_qos.cfg @@ -0,0 +1,37 @@ +SPECIFICATION Spec + +CONSTANTS + N = 3 + K = 2 + MaxConns = 3 + RMSet = {1, 2} + QSet = {1, 2} + RSet = {FALSE} + BSet = {FALSE} + MQSet = {0, 1, 2} + RASet = {TRUE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvRetainAvailable + InvMaximumQoS + InvMaximumPacketSize + InvReceiveMaximum + InvNoSilentLoss + InvRejectedNeverDelivered + InvOrder + InvPidUnique + InvPidNotQuarantined + InvNoStaleServerPid + InvNoMisattribution + InvExactlyOnce + InvRetainFidelity + InvQoSFidelity + InvFairActionEnabled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_rm2.cfg b/specs/tla/offline-queue/OfflineQueue_rm2.cfg new file mode 100644 index 00000000..bf70e689 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_rm2.cfg @@ -0,0 +1,37 @@ +SPECIFICATION Spec + +CONSTANTS + N = 3 + K = 2 + MaxConns = 3 + RMSet = {2} + QSet = {1, 2} + RSet = {TRUE, FALSE} + BSet = {FALSE} + MQSet = {2} + RASet = {TRUE, FALSE} + MBSet = {TRUE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvRetainAvailable + InvMaximumQoS + InvMaximumPacketSize + InvReceiveMaximum + InvNoSilentLoss + InvRejectedNeverDelivered + InvOrder + InvPidUnique + InvPidNotQuarantined + InvNoStaleServerPid + InvNoMisattribution + InvExactlyOnce + InvRetainFidelity + InvQoSFidelity + InvFairActionEnabled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/OfflineQueue_size.cfg b/specs/tla/offline-queue/OfflineQueue_size.cfg new file mode 100644 index 00000000..6dc24092 --- /dev/null +++ b/specs/tla/offline-queue/OfflineQueue_size.cfg @@ -0,0 +1,37 @@ +SPECIFICATION Spec + +CONSTANTS + N = 3 + K = 2 + MaxConns = 3 + RMSet = {1} + QSet = {1, 2} + RSet = {FALSE} + BSet = {TRUE, FALSE} + MQSet = {2} + RASet = {TRUE} + MBSet = {TRUE, FALSE} + Design = "Fix" + Quarantine = TRUE + Report = "handle" + ReplayDowngrade = FALSE + +INVARIANTS + TypeOK + InvRetainAvailable + InvMaximumQoS + InvMaximumPacketSize + InvReceiveMaximum + InvNoSilentLoss + InvRejectedNeverDelivered + InvOrder + InvPidUnique + InvPidNotQuarantined + InvNoStaleServerPid + InvNoMisattribution + InvExactlyOnce + InvRetainFidelity + InvQoSFidelity + InvFairActionEnabled + +CHECK_DEADLOCK FALSE diff --git a/specs/tla/offline-queue/README.md b/specs/tla/offline-queue/README.md new file mode 100644 index 00000000..64380963 --- /dev/null +++ b/specs/tla/offline-queue/README.md @@ -0,0 +1,294 @@ +# Offline queue and session resume: TLA+ model + +Model of the `mqtt5` client's outbound delivery across disconnects: the offline publish queue, +Session Present = 1 replay, Session Present = 0 recovery, and how each accepted QoS 1/2 publish +reports its outcome. It combines three independent models (a quorum) into one spec. The main +design is the one the user chose. The current code is kept as the variant `Design = "D0"` so +that the spec reproduces its bugs. + +## Files + +- `OfflineQueue.tla`: the spec, containing the chosen design (`Design = "Fix"`) and the current code (`Design = "D0"`). +- `OfflineQueue.cfg`, `OfflineQueue_rm2.cfg`, `OfflineQueue_qos.cfg`, `OfflineQueue_size.cfg`, `OfflineQueue_alldims.cfg`: main safety runs of the chosen design, split by dimension. +- `OfflineQueue_live.cfg`: liveness of the chosen design. +- `OfflineQueue_NEG_*.cfg`: negative controls, which must fail. +- `OfflineQueue_D0_*.cfg`: the current code, one cfg per property it breaks. +- `OfflineQueue_V_*.cfg`: rejected alternatives to the chosen design. Each one shows why the chosen mechanism is needed. + +## The chosen design + +1. **Per-publish outcome.** Every accepted QoS 1/2 publish gets its own completion handle. That + handle settles exactly once, to one of these outcomes: + - `Delivered{qos_used}` + - `Rejected(reason)`: the message was definitely not delivered. + - `Indeterminate(reason)`: the message may have been delivered. + + Outcomes are keyed by publish, never by packet id. The same rule covers live (online) + publishes and queued ones. If the connection drops while a live publish is in flight, its + handle stays pending and settles later. +2. **Enqueue-time checks.** A publish that breaks the last-known Retain Available or Maximum + Packet Size is rejected synchronously. Before the first CONNACK no limits are known, so this + check is permissive. A publish is also rejected synchronously when no packet id is free. +3. **Flush after replay, against the new CONNACK.** A queued message that no longer conforms + (RETAIN when Retain Available = 0, or larger than the new Maximum Packet Size) is handled like this: + - It is removed from the queue. + - Its packet id is freed. + - Its outcome is Rejected. If it was requeued after Session Present = 0 (so it may have been + delivered before), its outcome is Indeterminate instead. + - Flushing continues. There is no hold and no reordering. + + If the new Maximum QoS is lower than the requested QoS, the message is downgraded and sent. Its + outcome reports the QoS actually used. A downgrade to QoS 0 settles as `Delivered{0}` as soon + as it is written. `qos_used = 0` means unconfirmed. +4. **Session Present = 1 replay.** Each stored message is re-checked against the new CONNACK: + - A stored message in the PUBLISH stage that no longer conforms is not sent. Its outcome is + `Indeterminate(ReplayNotConforming)`. + - If that message was QoS 2, its packet id is **quarantined**: the id cannot be allocated + again until a Session Present = 0 connection. + - A stored QoS 2 message in the PUBREL stage is replayed as a PUBREL. PUBREL is not subject to + Retain Available, Maximum QoS or Maximum Packet Size. +5. **Session Present = 0.** + - Unacked messages that went out on the wire as QoS 1 return to the head of the queue in their + original order. They are then sent as new publishes: DUP=0, and new ids are allowed. + - Unacked QoS 2 messages in the PUBLISH stage get outcome Indeterminate. + - QoS 2 messages in the PUBREL stage get outcome `Delivered{2}`. Once the receiver has sent + PUBREC Success, it owns the message. + - Quarantine is cleared. +6. **Race-freedom.** + - Replay and flush work is bound to its connection. When a new connection starts, the task + from the old connection is killed. + - The step that takes a queued message is atomic: it removes the message from the queue by + identity, puts it in the session store and writes it. That step runs only on the current + connection. + - Packet ids of messages being conformed, stored or quarantined are excluded from allocation. +7. **Receive Maximum and ordering.** + - Receive Maximum is respected throughout, including during replay. A QoS 2 message holds its + slot until PUBCOMP. + - First receipts at the server follow publish order: replayed messages first, then queued + ones, then new live publishes. + +### Clarifications the design needed + +- **Replay downgrade (item 4).** A stored message whose QoS is above the new Maximum QoS is + treated as non-conforming. Its outcome is Indeterminate, and a QoS 2 id is quarantined. It is + not downgraded. The reasons: + - [MQTT-4.4.0-1] requires unacked PUBLISH packets to be resent with their original packet + identifiers, with DUP=1. The server may still hold QoS 2 state for that id: it received the + PUBLISH, but its PUBREC was lost. A QoS 1 PUBLISH under the same id would reuse an id that is + still in use by an unfinished QoS 2 exchange (a packet id stays in use until PUBCOMP). The + server would deliver the message again through its QoS 1 path and leave the QoS 2 state + orphaned. + - A QoS 0 resend is not a retransmission at all. QoS 0 forbids DUP=1 [MQTT-3.3.1-2] and has no + packet id. + - Sending above the new Maximum QoS is forbidden [MQTT-3.2.2-11]. + - So no spec-legal downgrade of an already-sent packet exists. The rejected alternative is + `ReplayDowngrade = TRUE` (`OfflineQueue_V_replaydowngrade.cfg`). The checker finds the + orphaned server QoS 2 state right away (`InvNoStaleServerPid`). +- **Downgrade and ExactlyOnce.** `qos_used` is the *lowest* QoS the message was ever sent at + (`effQ`). A message requeued after Session Present = 0 is re-flushed at + `Min(effQ, new Maximum QoS)`. It is never raised back to QoS 2, because a message that may + already have been delivered through the QoS 1 path cannot be made exactly-once afterwards. The + exactly-once guarantee applies only to messages sent at QoS 2 on every transmission. A message + received more than once must have `effQ < 2`, and any Delivered or Indeterminate outcome must + report `qos_used < 2`. +- **"Unacked QoS 1" at Session Present = 0** means QoS 1 *on the wire*. A QoS 2 request that was + downgraded to QoS 1 is therefore requeued. +- **PUBREL-stage replay** takes a Receive Maximum slot, because the message stays + unacknowledged until PUBCOMP. Replayed PUBRELs therefore never push the window over a smaller + new Receive Maximum. + +## Model + +- **Messages.** Messages `1..N` are published in order. Each has a requested QoS (`reqQ`), a + RETAIN flag (`reqR`) and an "oversize for a restrictive Maximum Packet Size" flag (`big`), + all chosen in `Init`. +- **Connections.** Each connection gets a CONNACK with Maximum QoS, Retain Available, + Maximum Packet Size, Receive Maximum and Session Present. Up to `MaxConns` connections happen. + Only a connection before the last one can be lost, so the last connection stays up. +- **Network.** The network is FIFO. A lost connection drops everything in flight. +- **Server.** The server delivers a PUBLISH when it receives it. It de-duplicates QoS 2 by + packet id (`srvQ2`) until it receives PUBREL. +- **QoS 2 handshake.** It is modelled as three steps: PUBREC, then the client's PUBREL, then + PUBCOMP. +- **Reporting.** `outcome[m]` and `outQ[m]` are the caller's view of message `m`. With + `Report = "pidEvent"` (a rejected alternative), results are instead emitted as events keyed by + packet id. The caller maps each event to whichever publish it last received that id for. + +## Properties in plain words + +All of these are invariants checked in every run of the chosen design. + +| Property | Meaning | +|---|---| +| `TypeOK` | Variables stay within their bounded types. | +| `InvRetainAvailable` | No PUBLISH on the wire (including replay) sets RETAIN while the current CONNACK says Retain Available = 0. | +| `InvMaximumQoS` | No PUBLISH on the wire exceeds the current Maximum QoS. | +| `InvMaximumPacketSize` | No PUBLISH on the wire exceeds the current Maximum Packet Size. | +| `InvReceiveMaximum` | The client never has more unacknowledged QoS 1/2 messages than the current Receive Maximum. | +| `InvNoSilentLoss` | An accepted message is in one of these situations: still held by the client, received by the server, or in flight. Otherwise its outcome says Rejected, Indeterminate or `Delivered{0}` (unconfirmed). "Delivered" at QoS > 0 for a message the server never received is a violation. This also catches a message swallowed by server de-duplication under a reused id. | +| `InvRejectedNeverDelivered` | A Rejected message was never received and is not in flight. | +| `InvOrder` | First receipts at the server are in publish order. | +| `InvPidUnique` | No two messages that the client still holds (queued, being conformed or popped, or stored) share a packet id. | +| `InvPidNotQuarantined` | No held message uses a quarantined id. | +| `InvNoStaleServerPid` | A packet on the wire never uses an id for which the server holds QoS 2 state belonging to a different message (or uses it at a QoS other than 2). No held message owns such an id. | +| `InvNoMisattribution` | An outcome is never delivered to the wrong publish. | +| `InvExactlyOnce` | A message received more than once was sent at QoS < 2 at least once, and its outcome reports `qos_used < 2`. Messages sent only at QoS 2 are received at most once. | +| `InvRetainFidelity` | RETAIN on the wire always equals what was requested. It is never silently cleared. No modelled design clears RETAIN, so this is a regression guard. | +| `InvQoSFidelity` | A Delivered or Indeterminate outcome reports exactly the lowest QoS used on the wire, so every downgrade is reported. | +| `InvFairActionEnabled` | While some accepted message is unresolved, one of the fair actions is enabled: connect, task step, server receive, client ack or caller consume. This progress check is based on `ENABLED`. It backs up the liveness results. | + +Liveness properties are checked under weak fairness of connect, task steps, server receive, +client ack and caller consume. + +| Property | Meaning | +|---|---| +| `LiveResolved == []<>AllResolved` | Every accepted message eventually gets its outcome. | +| `LiveSettled == []<>Settled` | The queue and the session store eventually drain. There is no head-of-line blocking and no stranded stored message. | + +Negative controls. Each must fail. + +| Property | Why it must fail | +|---|---| +| `NegImpossible == []<>(nextM > N + 1)` | Unreachable. | +| `NegAllDelivered` | A message can legitimately end Indeterminate or Rejected. | +| `NegQuarantineReleased == []<>(quar = {})` | If the last connection has Session Present = 1, the quarantine is held forever by design. | +| `LiveResolved` under `SpecNoConnectFairness` | Without fairness on Connect the client may stay offline forever. | + +## tla-mcp 0.9.4 caveat + +The modelers found two liveness defects in tla-mcp 0.9.4: + +- `P ~> Q` can pass vacuously. +- Missing fairness is not honoured: the checker behaves as if strong fairness were present. + +To work around them: + +- All liveness is written as `[]<>` over predicates that become stable. Nothing uses `~>`. +- Negative controls are included, and their results are recorded below. +- Every safety run includes `InvFairActionEnabled`. + +`OfflineQueue_NEG_noconnectfair.cfg` is the probe for the fairness defect. Without fairness on +Connect, `LiveResolved` must fail. The tool reports `ok`, which reproduces the defect. Because of +this, the positive liveness results are evidence, not proof: they are only as strong as the +`[]<>` form, the negative controls and `InvFairActionEnabled` together. + +A second tooling limit: the MCP client stops a `check_spec` call that is silent for 1800 s. Any +run that needed longer is listed as inconclusive and was split by dimension instead. + +## Runs of the chosen design (final spec) + +All runs use `K = 2` packet ids, `Design = "Fix"`, `Quarantine = TRUE`, `Report = "handle"` and +`ReplayDowngrade = FALSE`. Every run checks all 16 invariants listed above. Some runs shared the +machine, so the elapsed times include contention. + +| cfg | N | MaxConns | RMSet | QSet | RSet | BSet | MQSet | RASet | MBSet | Result | States | Depth | Secs | +|---|---|---|---|---|---|---|---|---|---|---|---|---|---| +| `OfflineQueue.cfg` (RETAIN / Retain Available dimension) | 3 | 3 | {1} | {1,2} | {T,F} | {F} | {2} | {T,F} | {T} | ok | 339,852 | 31 | 450 | +| `OfflineQueue_rm2.cfg` | 3 | 3 | {2} | {1,2} | {T,F} | {F} | {2} | {T,F} | {T} | ok | 589,643 | 33 | 799 | +| `OfflineQueue_qos.cfg` (Maximum QoS dimension, including downgrade to 0) | 3 | 3 | {1,2} | {1,2} | {F} | {F} | {0,1,2} | {T} | {T} | ok | 402,577 | 33 | 660 | +| `OfflineQueue_size.cfg` (Maximum Packet Size dimension) | 3 | 3 | {1} | {1,2} | {F} | {T,F} | {2} | {T} | {T,F} | ok | 339,852 | 31 | 468 | +| `OfflineQueue_alldims.cfg` (every dimension at once) | 2 | 2 | {1,2} | {1,2} | {T,F} | {T,F} | {0,1,2} | {T,F} | {T,F} | ok | 751,999 | 22 | 1323 | +| `OfflineQueue_live.cfg` (liveness: `LiveResolved`, `LiveSettled`) | 2 | 3 | {1,2} | {1,2} | {T,F} | {F} | {0,2} | {T,F} | {T} | ok | 157,270 | 26 | 303 | + +These runs were not completed and are recorded as inconclusive, not as passes: + +- The RETAIN dimension with `RMSet = {1,2}` at N=3, MaxConns=3 hit `limit_reached` at 935,089 + states (depth 23) after 1700 s, because of the 1800 s MCP limit. It is covered by the two runs + split by Receive Maximum, plus `OfflineQueue_qos.cfg`, which mixes Receive Maximum 1 and 2 + across connections. +- An all-dimensions run at N=2, MaxConns=3 on an earlier revision of the spec was stopped by the + MCP timeout. +- The Maximum Packet Size dimension with `RMSet = {2}` was not run. The chosen design treats + Maximum Packet Size and Retain Available non-conformance through identical code paths, and the + two `RMSet = {1}` runs produced the same state graph (339,852 states each). + +## Current code (D0), variants and negative controls + +The run results for this section are listed at the end of this file, in the section "D0, +variant and negative-control results (final spec)". + +What D0 models. It follows the code at the time of modelling (`client/direct/replay.rs`, +`client/direct/mod.rs`): + +- An offline publish returns `Ok(packet_id)` at once. It has no completion, so the caller + believes it will be delivered. +- A flush that finds a non-conforming message drops it silently. +- A Maximum QoS downgrade is silent. +- Replay resends stored messages as they are, against the new CONNACK. +- Session Present = 0 discards the session state. Only live-publish futures learn of it, and + they get Indeterminate. +- The replay task is tied to its connection only through a `Weak` writer, and it pops the + shared queue by position (`pop_front`). +- `pop_front` runs before `store_unacked_publish`, which leaves a gap in which the packet id + belongs neither to the queue nor to the store. The allocator does not see ids in that gap. + +## Mapping from spec actions to client code + +The names below are as of this modelling pass. The code is being reworked alongside. + +| Spec | Client concept | +|---|---| +| `Accept` | `stage_publish` / `queue_publish_message`, with enqueue-time checks (`check_publish_size` plus a Retain Available check) and `allocate_packet_id`. In the chosen design the allocator also excludes quarantined ids and ids of messages being conformed. | +| `outcome`, `outQ`, `Resolve` (`Report = "handle"`) | The per-publish completion handle: Delivered{qos_used}, Rejected or Indeterminate. Today online publishes use `PublishAck` futures, and queued ones get `PublishResult::QoS1Or2 { packet_id }` with no completion. | +| `Report = "pidEvent"`, `CallerConsume` | The rejected alternative: outcomes reported through a packet-id-keyed callback or event stream. | +| `Connect` | `connect`: `apply_server_capabilities` (Session Present = 0 leads to `discard_session_state`), `reset_send_quota`, `advance_connection_epoch`, then spawning `SessionReplay`. In the chosen design Session Present = 0 requeues QoS 1 at the head of `OfflineQueue` and resolves QoS 2 as Indeterminate (PUBLISH stage) or Delivered (PUBREL stage). Connect also kills the previous connection's task and clears quarantine when Session Present = 0. | +| `Lose` | Transport or reader failure. | +| `ReplayStep`, `SendReplay`, `SendReplayRel`, `AbandonReplay` | `SessionReplay::replay_session_state` (`OutboundReplay::Publish` / `PubRel`), with a re-check of each stored PUBLISH against the new CONNACK. `AbandonReplay` resolves Indeterminate and quarantines a QoS 2 id. | +| `Conform`, `Drop`, `Claim` (`ClaimFix`) | `SessionReplay::flush_offline_queue`, `conform_queued` and `take_slot`. `ClaimFix` is the required atomic step: check the connection epoch, remove the message by identity, store it, then write it. | +| `ClaimD0`, `StoreStep`, `WriteStep` | Today's `pop_front`, then `store_unacked_publish`, then `write`. These steps are separate and are not bound to the connection. | +| `task[e]`, `e = conns` guards | The connection epoch that replay and flush work is bound to. | +| `quar` | New state: QoS 2 ids abandoned during replay, excluded from allocation until Session Present = 0. | +| `infl`, `caps.rm` | The `FlowControlManager` send quota (`claim_send_quota`, `acknowledge`). | +| `ServerRecv`, `ClientAck` | The broker, plus the client's PUBACK/PUBREC/PUBCOMP handling (`ack.rs`, `handlers.rs`, `release_outbound_quota`). | + +## Abstractions and refinements outside the model + +- **Live publishes.** The model sends every accepted publish, online or offline, through the same + queue, per-publish outcome and flush step. A publish made while connected, with the flush idle, + goes out immediately through that path. So the late design update ("live publishes get the + same per-publish outcome; connection loss leaves the handle pending") is what the model + already checks. `NoSilentLoss`, `LiveResolved` and the exactly-once resolution of `outcome` + cover live publishes too. +- **QoS 2 PUBREL stage.** QoS 2 is split into PUBREC, PUBREL and PUBCOMP. A PUBREL-stage message + at Session Present = 0 resolves Delivered{2}, which is the second late design update. At + Session Present = 1 its PUBREL is replayed. +- **Server behaviour.** The server always accepts: there are no error reason codes. Two + implementation refinements are therefore outside the model: + - The mismatched-ack protocol-error rule: an ack whose type does not match the stored exchange + is a protocol error. + - "An error PUBREC removes state before reporting Rejected." + + Both refine how a server refusal ends an exchange. Neither adds a new way to deliver or lose a + message. +- **RETAIN, oversize and Maximum QoS** are abstracted to one flag or level per message. Topic + aliases, the payload and the DUP bit are not modelled. DUP is implied: a replay resends the + same id. +- **Bounds.** The results are bounded model checking. They hold for the constants listed, not in + general. + +## D0, variant and negative-control results (final spec) + +Each safety cfg checks `TypeOK` and one property. BFS returns the shortest counterexample. + +| cfg | Expected | Result | States | Depth | What the counterexample shows | +|---|---|---|---|---|---| +| `OfflineQueue_D0_NoSilentLoss.cfg` | violation | `InvNoSilentLoss` violated | 457 | 6 | An offline publish is told Ok. After reconnecting with a smaller Maximum Packet Size, the flush drops it with no report. | +| `OfflineQueue_D0_RetainAvailable.cfg` | violation | `InvRetainAvailable` violated | 741 | 10 | A stored RETAIN message is replayed after a reconnect with Retain Available = 0. | +| `OfflineQueue_D0_MaximumPacketSize.cfg` | violation | `InvMaximumPacketSize` violated | 684 | 10 | A stored message is replayed after a reconnect with a smaller Maximum Packet Size. | +| `OfflineQueue_D0_MaximumQoS.cfg` | violation | `InvMaximumQoS` violated | 788 | 10 | A stored QoS 2 message is replayed at QoS 2 after a reconnect with Maximum QoS 1. | +| `OfflineQueue_D0_QoSFidelity.cfg` | violation | `InvQoSFidelity` violated | 289 | 8 | A QoS 2 publish is silently sent at QoS 1, while the caller was told Ok at QoS 2. | +| `OfflineQueue_D0_PidUnique.cfg` | violation | `InvPidUnique` violated | 70 | 7 | In the pop/store gap the allocator hands a new publish an id that is still in flight. | +| `OfflineQueue_D0_ExactlyOnce.cfg` | violation | `InvExactlyOnce` violated | 2,662 | 13 | A requested QoS 2 message is silently downgraded to QoS 1 and replayed, so it is delivered twice while reported as QoS 2. | +| `OfflineQueue_D0_Order.cfg` (MaxConns = 3) | violation | `InvOrder` violated | 21,739 | 20 | A stale task pops m1, but m1 is stored only after the next connection's replay snapshot. m2 is delivered first, and m1 arrives on the following reconnect. | +| `OfflineQueue_D0_NoStaleServerPid.cfg` | violation | `InvNoStaleServerPid` violated | 261 | 10 | The pop/store gap reuses the id of an in-flight QoS 2 message. The server would de-duplicate the new message away. | +| `OfflineQueue_D0_live.cfg` | liveness failure | `LiveSettled` violated | 37,703 | 30 | A message stored by a stale task misses the replay snapshot and stays in the store forever. | +| `OfflineQueue_V_noquarantine.cfg` | violation | `InvNoStaleServerPid` violated | 1,021 | 11 | Without quarantine, the id of a QoS 2 message abandoned during replay is reused while the server still holds its QoS 2 state. | +| `OfflineQueue_V_replaydowngrade.cfg` | violation | `InvNoStaleServerPid` violated | 200 | 10 | A replay downgrades a stored QoS 2 message to QoS 1 under the same id while the server holds QoS 2 state for it. | +| `OfflineQueue_V_pidevent.cfg` | violation | `InvNoMisattribution` violated | 409 | 8 | ABA misattribution: the Rejected event for id 1 (m1) is consumed after id 1 was reallocated to m2, so m2 is reported Rejected. | +| `OfflineQueue_NEG_impossible.cfg` | liveness failure | fails | 157,270 | 26 | The unreachable goal is correctly reported. | +| `OfflineQueue_NEG_alldelivered.cfg` | liveness failure | fails | 157,270 | 26 | A message ends Indeterminate. | +| `OfflineQueue_NEG_quarantine.cfg` | liveness failure | fails | 157,270 | 26 | A QoS 2 id is quarantined, and the final connection has Session Present = 1. | +| `OfflineQueue_NEG_noconnectfair.cfg` | liveness failure | **ok (tool defect)** | 157,270 | 26 | The missing fairness is not honoured, which reproduces the known tla-mcp 0.9.4 defect. | + +The liveness runs of the chosen design (`OfflineQueue_live.cfg`) and all the negative controls +use the same constants.