use anyhow::{anyhow, bail, Context, Result}; use serde::{Deserialize, Serialize}; use std::{ sync::{ atomic::{AtomicBool, AtomicU16, AtomicU64, Ordering}, Arc, }, time::{Duration, SystemTime, UNIX_EPOCH}, }; use tokio::{ io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}, net::TcpStream, sync::{broadcast, mpsc}, time::{interval, timeout}, }; use tokio_rustls::{rustls, TlsConnector}; use tokio_rustls::rustls::pki_types::ServerName; pub const MQTT_PORT: u16 = 1984; pub const MQTT_KEEPALIVE_SECONDS: u16 = 60; const MQTT_QUEUE_TIMEOUT: Duration = Duration::from_secs(2); const MQTT_WRITE_TIMEOUT: Duration = Duration::from_secs(5); const MQTT_DISCONNECT_TIMEOUT: Duration = Duration::from_millis(750); pub fn broker_for_region(region: &str) -> Option<&'static str> { match region { "Australia" => Some("mqtt-au.gree.com"), // greeclimate 1.2.1 regional broker mapping. "China Mainland" => Some("mqtt-cn.gree.com"), "East South Asia" => Some("mqtt-as.gree.com"), "Europe" => Some("mqtt-eu.gree.com"), "India" => Some("mqtt-in.gree.com"), "Latin American" => Some("mqtt-la.gree.com"), "Middle East" => Some("mqtt-me.gree.com"), "North American" => Some("mqtt-na.gree.com"), "Russia" => Some("mqtt-ru.gree.com"), "South American" => Some("mqtt-sa.gree.com"), _ => None, } } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MqttDeviceEnvelope { #[serde(default)] pub cid: String, #[serde(default)] pub i: i64, #[serde(default)] pub pack: String, #[serde(default)] pub t: String, #[serde(default)] pub tcid: String, #[serde(default)] pub uid: serde_json::Value, #[serde(default, skip_serializing_if = "Option::is_none")] pub tag: Option, #[serde(default, skip_serializing_if = "Option::is_none")] pub ts: Option, } #[derive(Debug, Clone)] pub enum MqttEvent { Message { topic: String, payload: Vec }, /// Broker-level traffic such as SUBACK/PUBACK/PINGRESP. This proves the MQTT /// session is alive without being mistaken for a response from the HVAC unit. Traffic { kind: &'static str }, Disconnected { reason: String }, } #[derive(Debug)] enum WireCommand { Subscribe(Vec), Publish { topic: String, payload: Vec }, Raw(Vec), Ping, Disconnect, } #[derive(Clone)] pub struct MqttConnection { tx: mpsc::Sender, connected: Arc, } impl MqttConnection { pub async fn connect( host: &str, port: u16, user_id: i64, token: &str, connect_timeout: Duration, events: broadcast::Sender, ) -> Result { let tcp = timeout(connect_timeout, TcpStream::connect((host, port))) .await .context("GREE Cloud MQTT TCP connect timed out")? .with_context(|| format!("cannot connect to GREE Cloud MQTT {host}:{port}"))?; tcp.set_nodelay(true).ok(); // Do not inherit the reference library's CERT_NONE workaround: certificate and // hostname verification are intentionally enabled here. let mut roots = rustls::RootCertStore::empty(); roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); let tls_config = rustls::ClientConfig::builder() .with_root_certificates(roots) .with_no_client_auth(); let connector = TlsConnector::from(Arc::new(tls_config)); let server_name = ServerName::try_from(host.to_string()) .map_err(|_| anyhow!("invalid GREE Cloud MQTT hostname"))?; let mut stream = timeout(connect_timeout, connector.connect(server_name, tcp)) .await .context("GREE Cloud MQTT TLS handshake timed out")? .context("GREE Cloud MQTT TLS handshake failed")?; let client_id = format!("app_{:016x}", rand::random::()); let connect_packet = encode_connect( &client_id, &user_id.to_string(), token, MQTT_KEEPALIVE_SECONDS, )?; timeout(connect_timeout, stream.write_all(&connect_packet)) .await .context("GREE Cloud MQTT CONNECT write timed out")??; timeout(connect_timeout, stream.flush()) .await .context("GREE Cloud MQTT CONNECT flush timed out")??; let (packet_type, payload) = timeout(connect_timeout, read_packet(&mut stream)) .await .context("GREE Cloud MQTT CONNACK timed out")??; if packet_type >> 4 != 2 || payload.len() != 2 { bail!("invalid GREE Cloud MQTT CONNACK"); } if payload[1] != 0 { let reason = match payload[1] { 1 => "unacceptable protocol version", 2 => "identifier rejected", 3 => "server unavailable", 4 => "invalid username/password", 5 => "not authorized", _ => "unknown broker error", }; bail!("GREE Cloud MQTT authentication/connect rejected: {reason}"); } let (reader, writer) = tokio::io::split(stream); let (tx, rx) = mpsc::channel(64); let connected = Arc::new(AtomicBool::new(true)); let packet_ids = Arc::new(AtomicU16::new(1)); let last_rx_ms = Arc::new(AtomicU64::new(unix_millis())); spawn_writer( writer, rx, connected.clone(), events.clone(), packet_ids, ); spawn_reader( reader, tx.clone(), connected.clone(), events.clone(), last_rx_ms.clone(), ); spawn_keepalive(tx.clone(), connected.clone(), events.clone(), last_rx_ms); Ok(Self { tx, connected, }) } pub fn is_connected(&self) -> bool { self.connected.load(Ordering::Acquire) } pub async fn subscribe_device(&self, parent_mac: &str) -> Result<()> { if !self.is_connected() { bail!("GREE Cloud MQTT is not connected"); } let topics = [ format!("response/{parent_mac}/#"), format!("status/{parent_mac}/#"), format!("connect/{parent_mac}"), ]; // Match the reference client: one QoS1 SUBSCRIBE per topic. This is slightly more // verbose than a multi-filter packet but avoids broker-specific handling differences. for topic in topics { timeout( MQTT_QUEUE_TIMEOUT, self.tx.send(WireCommand::Subscribe(vec![topic])), ) .await .context("GREE Cloud MQTT subscribe queue timed out")? .map_err(|_| anyhow!("GREE Cloud MQTT writer stopped"))?; } Ok(()) } pub async fn publish(&self, topic: String, payload: Vec) -> Result<()> { if !self.is_connected() { bail!("GREE Cloud MQTT is not connected"); } timeout( MQTT_QUEUE_TIMEOUT, self.tx.send(WireCommand::Publish { topic, payload }), ) .await .context("GREE Cloud MQTT publish queue timed out")? .map_err(|_| anyhow!("GREE Cloud MQTT writer stopped")) } pub async fn disconnect(&self) { // Mark disconnected first so no new work can enter the queue while shutdown is in // progress. A wedged/full writer queue must never prevent process termination. self.connected.store(false, Ordering::Release); let _ = timeout( MQTT_DISCONNECT_TIMEOUT, self.tx.send(WireCommand::Disconnect), ) .await; } } fn spawn_keepalive( tx: mpsc::Sender, connected: Arc, events: broadcast::Sender, last_rx_ms: Arc, ) { tokio::spawn(async move { let mut tick = interval(Duration::from_secs(30)); tick.tick().await; while connected.load(Ordering::Acquire) { tick.tick().await; let idle_ms = unix_millis().saturating_sub(last_rx_ms.load(Ordering::Acquire)); if idle_ms > 90_000 { mark_disconnected( &connected, &events, "MQTT heartbeat timed out waiting for broker traffic".into(), ); let _ = timeout(MQTT_DISCONNECT_TIMEOUT, tx.send(WireCommand::Disconnect)).await; break; } match timeout(MQTT_QUEUE_TIMEOUT, tx.send(WireCommand::Ping)).await { Ok(Ok(())) => {} Ok(Err(_)) => { mark_disconnected(&connected, &events, "MQTT writer stopped".into()); break; } Err(_) => { mark_disconnected(&connected, &events, "MQTT ping queue timed out".into()); break; } } } }); } fn spawn_writer( mut writer: W, mut rx: mpsc::Receiver, connected: Arc, events: broadcast::Sender, packet_ids: Arc, ) where W: AsyncWrite + Unpin + Send + 'static, { tokio::spawn(async move { while let Some(command) = rx.recv().await { let result = match command { WireCommand::Subscribe(topics) => { let packet_id = next_packet_id(&packet_ids); encode_subscribe(packet_id, &topics) } WireCommand::Publish { topic, payload } => { let packet_id = next_packet_id(&packet_ids); encode_publish(packet_id, &topic, &payload) } WireCommand::Raw(packet) => Ok(packet), WireCommand::Ping => Ok(vec![0xC0, 0x00]), WireCommand::Disconnect => { let _ = timeout(MQTT_DISCONNECT_TIMEOUT, writer.write_all(&[0xE0, 0x00])).await; let _ = timeout(MQTT_DISCONNECT_TIMEOUT, writer.flush()).await; connected.store(false, Ordering::Release); break; } }; let packet = match result { Ok(packet) => packet, Err(err) => { mark_disconnected(&connected, &events, err.to_string()); break; } }; match timeout(MQTT_WRITE_TIMEOUT, writer.write_all(&packet)).await { Ok(Ok(())) => {} Ok(Err(err)) => { mark_disconnected(&connected, &events, format!("MQTT write failed: {err}")); break; } Err(_) => { mark_disconnected(&connected, &events, "MQTT write timed out".into()); break; } } match timeout(MQTT_WRITE_TIMEOUT, writer.flush()).await { Ok(Ok(())) => {} Ok(Err(err)) => { mark_disconnected(&connected, &events, format!("MQTT flush failed: {err}")); break; } Err(_) => { mark_disconnected(&connected, &events, "MQTT flush timed out".into()); break; } } } }); } fn spawn_reader( mut reader: R, tx: mpsc::Sender, connected: Arc, events: broadcast::Sender, last_rx_ms: Arc, ) where R: AsyncRead + Unpin + Send + 'static, { tokio::spawn(async move { loop { match read_packet(&mut reader).await { Ok((header, payload)) => { last_rx_ms.store(unix_millis(), Ordering::Release); match header >> 4 { 3 => { if let Err(err) = handle_publish(header, &payload, &tx, &events).await { tracing::warn!(error=?err, "invalid GREE Cloud MQTT PUBLISH"); } } 9 => { let _ = events.send(MqttEvent::Traffic { kind: "SUBACK" }); } 4 => { let _ = events.send(MqttEvent::Traffic { kind: "PUBACK" }); } 13 => { let _ = events.send(MqttEvent::Traffic { kind: "PINGRESP" }); } _ => {} } } Err(err) => { mark_disconnected(&connected, &events, format!("MQTT read failed: {err}")); break; } } } }); } async fn handle_publish( header: u8, payload: &[u8], tx: &mpsc::Sender, events: &broadcast::Sender, ) -> Result<()> { if payload.len() < 2 { bail!("short MQTT publish packet"); } let topic_len = u16::from_be_bytes([payload[0], payload[1]]) as usize; if payload.len() < 2 + topic_len { bail!("truncated MQTT publish topic"); } let topic = std::str::from_utf8(&payload[2..2 + topic_len])?.to_string(); let qos = (header >> 1) & 0x03; let mut offset = 2 + topic_len; if qos > 0 { if payload.len() < offset + 2 { bail!("truncated MQTT publish packet id"); } let packet_id = u16::from_be_bytes([payload[offset], payload[offset + 1]]); offset += 2; if qos == 1 { // MQTT QoS1 requires a PUBACK for every incoming publish. Send the raw four-byte // acknowledgement through the single writer task so frame writes never interleave. let ack = vec![0x40, 0x02, (packet_id >> 8) as u8, packet_id as u8]; send_raw_ack(tx, ack).await?; } } let body = payload[offset..].to_vec(); let _ = events.send(MqttEvent::Message { topic, payload: body }); Ok(()) } // MQTT QoS1 delivery requires PUBACK. To keep WireCommand's public operations minimal, encode // acknowledgement as a synthetic command handled by a reserved topic marker. async fn send_raw_ack(tx: &mpsc::Sender, ack: Vec) -> Result<()> { timeout(MQTT_QUEUE_TIMEOUT, tx.send(WireCommand::Raw(ack))) .await .context("GREE Cloud MQTT ACK queue timed out")? .map_err(|_| anyhow!("GREE Cloud MQTT writer stopped")) } fn unix_millis() -> u64 { SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap_or_default() .as_millis() .min(u128::from(u64::MAX)) as u64 } fn next_packet_id(ids: &AtomicU16) -> u16 { let id = ids.fetch_add(1, Ordering::Relaxed); if id == 0 { 1 } else { id } } fn mark_disconnected( connected: &AtomicBool, events: &broadcast::Sender, reason: String, ) { if connected.swap(false, Ordering::AcqRel) { let _ = events.send(MqttEvent::Disconnected { reason }); } } fn encode_connect(client_id: &str, username: &str, password: &str, keepalive: u16) -> Result> { let mut body = Vec::new(); push_utf8(&mut body, "MQTT")?; body.push(4); // MQTT 3.1.1 body.push(0xC2); // username + password + clean session body.extend_from_slice(&keepalive.to_be_bytes()); push_utf8(&mut body, client_id)?; push_utf8(&mut body, username)?; push_utf8(&mut body, password)?; frame(0x10, body) } fn encode_subscribe(packet_id: u16, topics: &[String]) -> Result> { let mut body = Vec::new(); body.extend_from_slice(&packet_id.to_be_bytes()); for topic in topics { push_utf8(&mut body, topic)?; body.push(1); // requested QoS 1 } frame(0x82, body) } fn encode_publish(packet_id: u16, topic: &str, payload: &[u8]) -> Result> { let mut body = Vec::new(); push_utf8(&mut body, topic)?; body.extend_from_slice(&packet_id.to_be_bytes()); body.extend_from_slice(payload); frame(0x32, body) // PUBLISH QoS1 } fn frame(header: u8, body: Vec) -> Result> { let mut out = Vec::with_capacity(body.len() + 5); out.push(header); encode_remaining_length(body.len(), &mut out)?; out.extend_from_slice(&body); Ok(out) } fn push_utf8(out: &mut Vec, value: &str) -> Result<()> { let len = u16::try_from(value.as_bytes().len()).context("MQTT string is too long")?; out.extend_from_slice(&len.to_be_bytes()); out.extend_from_slice(value.as_bytes()); Ok(()) } fn encode_remaining_length(mut len: usize, out: &mut Vec) -> Result<()> { if len > 268_435_455 { bail!("MQTT packet is too large"); } loop { let mut digit = (len % 128) as u8; len /= 128; if len > 0 { digit |= 0x80; } out.push(digit); if len == 0 { break; } } Ok(()) } async fn read_packet(reader: &mut R) -> Result<(u8, Vec)> { let header = reader.read_u8().await?; let mut multiplier = 1usize; let mut remaining = 0usize; for _ in 0..4 { let digit = reader.read_u8().await?; remaining = remaining .checked_add(((digit & 0x7f) as usize).saturating_mul(multiplier)) .ok_or_else(|| anyhow!("invalid MQTT remaining length"))?; if digit & 0x80 == 0 { let mut payload = vec![0_u8; remaining]; reader.read_exact(&mut payload).await?; return Ok((header, payload)); } multiplier = multiplier.saturating_mul(128); } bail!("invalid MQTT remaining length encoding") } #[cfg(test)] mod tests { use super::*; #[test] fn mqtt_connect_uses_v311_and_credentials() { let packet = encode_connect("app_123", "42", "secret", 60).unwrap(); assert_eq!(packet[0], 0x10); assert!(packet.windows(6).any(|w| w == b"\0\x04MQTT")); assert!(packet.windows(4).any(|w| w == b"\0\x0242")); } #[test] fn region_brokers_match_cloud_reference() { assert_eq!(broker_for_region("Europe"), Some("mqtt-eu.gree.com")); assert_eq!(broker_for_region("North American"), Some("mqtt-na.gree.com")); assert_eq!(broker_for_region("China Mainland"), Some("mqtt-cn.gree.com")); } }