Reinforcement Learning
stable-baselines3
deep-reinforcement-learning
agricultural-ai
weather-modelling
curriculum-learning
edge-ai
Instructions to use DHDRL/monsoon-rl with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- stable-baselines3
How to use DHDRL/monsoon-rl with stable-baselines3:
from huggingface_sb3 import load_from_hub checkpoint = load_from_hub( repo_id="DHDRL/monsoon-rl", filename="{MODEL FILENAME}.zip", ) - Notebooks
- Google Colab
- Kaggle
| """ | |
| node_transport.py | |
| ================= | |
| Transport-neutral node bus for edge fleet coordination. | |
| Semantics (all backends): | |
| - Latest-value, non-consuming reads | |
| - Soft TTL (STATE_TTL_SEC); recv returns None when stale | |
| - No pickle / no eval — versioned binary or JSON only | |
| Wire envelope (both hidden state and product alerts): | |
| [version:u8][flags:u8][payload...] | |
| flags bit0 = zlib-compressed payload | |
| Hidden inner payload (unchanged for torch side): | |
| [n_layers:u32][batch:u32][hidden:u32][float32 data...] little-endian | |
| Alert inner payload: | |
| UTF-8 JSON object (see alert_to_bytes / bytes_to_alert) | |
| NodeTransport protocol (Genius-safe abstraction): | |
| send/recv_hidden_state, send/recv_alert, clear | |
| MQTT is one backend (hub-and-spoke). A future libp2p / SuperGenius backend | |
| implements the same protocol without callers knowing topic names or brokers. | |
| Steps covered here: | |
| 1. Product alerts on the protocol (not MQTT-only helpers) | |
| 2. MQTT hardening: TLS, auth, LWT + birth presence, reconnect | |
| 3. MQTT 5 when available (message expiry, user properties); 3.1.1 fallback | |
| 4. Schema version + optional zlib on all wire payloads | |
| 5. Multi-broker list (try in order) — topology as config, not hard-coded sites | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import logging | |
| import struct | |
| import threading | |
| import time | |
| import zlib | |
| from dataclasses import asdict, dataclass, field | |
| from typing import Any, Dict, List, Optional, Protocol, Sequence, Tuple, runtime_checkable | |
| import zone_observation as _zo | |
| assert _zo.SCHEMA_VERSION == 3, ( | |
| f"node_transport: zone_observation schema mismatch " | |
| f"(expected 3, got {_zo.SCHEMA_VERSION})" | |
| ) | |
| logger = logging.getLogger(__name__) | |
| # --------------------------------------------------------------------------- | |
| # Constants | |
| # --------------------------------------------------------------------------- | |
| MAX_HIDDEN_BYTES = 10 * 1024 * 1024 # 10 MB cap (pre-envelope) | |
| MAX_ALERT_BYTES = 64 * 1024 # 64 KB alert JSON cap | |
| STATE_TTL_SEC = 300 # soft TTL for latest-value reads | |
| WIRE_VERSION = 1 | |
| FLAG_COMPRESSED = 0x01 | |
| # MQTT topic layout is an implementation detail of MQTTTransport only. | |
| DEFAULT_TOPIC_ROOT = "weather" | |
| # weather/hidden/{zone_id} | |
| # weather/alert/{zone_id} | |
| # weather/presence/{node_id} | |
| # --------------------------------------------------------------------------- | |
| # Wire envelope (transport-neutral) | |
| # --------------------------------------------------------------------------- | |
| def pack_wire(payload: bytes, *, compress: bool = False) -> bytes: | |
| """Prefix payload with version + flags. Optional zlib on payload only.""" | |
| flags = 0 | |
| body = payload | |
| if compress and len(payload) > 64: | |
| body = zlib.compress(payload, level=6) | |
| flags |= FLAG_COMPRESSED | |
| return bytes([WIRE_VERSION, flags]) + body | |
| def unpack_wire(data: bytes) -> bytes: | |
| """ | |
| Strip version envelope. Accepts: | |
| - versioned: [ver][flags][payload] | |
| - legacy hidden: raw [u32][u32][u32][floats...] (no version byte) | |
| """ | |
| if len(data) < 2: | |
| raise ValueError(f"wire payload too short: {len(data)}") | |
| # Legacy hidden: first uint32 is n_layers in 1..8 and total length matches | |
| if _looks_like_legacy_hidden(data): | |
| return data | |
| ver, flags = data[0], data[1] | |
| if ver != WIRE_VERSION: | |
| raise ValueError(f"unsupported wire version: {ver}") | |
| body = data[2:] | |
| if flags & FLAG_COMPRESSED: | |
| try: | |
| body = zlib.decompress(body) | |
| except zlib.error as e: | |
| raise ValueError(f"zlib decompress failed: {e}") from e | |
| return body | |
| def _looks_like_legacy_hidden(data: bytes) -> bool: | |
| if len(data) < 12: | |
| return False | |
| n_layers, batch_size, hidden_size = struct.unpack_from("<III", data, 0) | |
| if not (1 <= n_layers <= 8 and 1 <= batch_size <= 4096 and 1 <= hidden_size <= 4096): | |
| return False | |
| expected = 12 + n_layers * batch_size * hidden_size * 4 | |
| return len(data) == expected | |
| # --------------------------------------------------------------------------- | |
| # Hidden-state serialisation (torch side — inner payload only) | |
| # --------------------------------------------------------------------------- | |
| def hidden_to_bytes(hidden_tensor) -> bytes: | |
| """ | |
| Serialise a GRU hidden state tensor to inner payload bytes. | |
| Format: | |
| [n_layers: uint32][batch_size: uint32][hidden_size: uint32][float32 data...] | |
| Always little-endian, always float32. | |
| Call pack_wire(...) before sending on the bus if compression/versioning needed; | |
| LocalTransport and MQTTTransport pack automatically on send. | |
| """ | |
| import torch | |
| h = hidden_tensor.detach().cpu().float() | |
| n_layers, batch_size, hidden_size = h.shape | |
| header = struct.pack("<III", n_layers, batch_size, hidden_size) | |
| return header + h.numpy().tobytes() | |
| def bytes_to_hidden(data: bytes): | |
| """ | |
| Safe deserialisation with strict validation before any ML import. | |
| Accepts wire envelope or legacy/inner payload. | |
| """ | |
| if _looks_like_legacy_hidden(data): | |
| inner = data | |
| else: | |
| inner = unpack_wire(data) | |
| if len(inner) < 12: | |
| raise ValueError(f"Hidden state bytes too short: {len(inner)}") | |
| n_layers, batch_size, hidden_size = struct.unpack_from("<III", inner, 0) | |
| if not (1 <= n_layers <= 8 and 1 <= batch_size <= 4096 and 1 <= hidden_size <= 4096): | |
| raise ValueError( | |
| f"Hidden state shape out of bounds: " | |
| f"({n_layers}, {batch_size}, {hidden_size})" | |
| ) | |
| expected_floats = n_layers * batch_size * hidden_size | |
| expected_bytes = 12 + expected_floats * 4 | |
| if len(inner) != expected_bytes: | |
| raise ValueError( | |
| f"Hidden state byte length mismatch: " | |
| f"got {len(inner)}, expected {expected_bytes}" | |
| ) | |
| import numpy as np | |
| import torch | |
| arr = np.frombuffer(inner[12:], dtype=np.float32).reshape( | |
| n_layers, batch_size, hidden_size | |
| ) | |
| return torch.from_numpy(arr.copy()) | |
| # --------------------------------------------------------------------------- | |
| # Product alert serialisation (transport-neutral) | |
| # --------------------------------------------------------------------------- | |
| class ProductAlert: | |
| """Client/product-facing alert snapshot for the node bus. | |
| Mirrors env terminate info + scorer freeze fields. Transport-agnostic: | |
| any NodeTransport backend can carry the same bytes. | |
| """ | |
| zone_id: str | |
| product_actionable: bool | |
| elevated: bool | |
| alert_level: str | |
| drought_risk: float = 0.0 | |
| flood_risk: float = 0.0 | |
| max_risk: float = 0.0 | |
| confidence: float = 0.0 | |
| trigger: str = "" | |
| ts: float = field(default_factory=time.time) | |
| schema_version: int = 1 | |
| extras: Dict[str, Any] = field(default_factory=dict) | |
| def to_dict(self) -> Dict[str, Any]: | |
| d = asdict(self) | |
| return d | |
| def alert_to_bytes(alert: ProductAlert) -> bytes: | |
| """JSON inner payload for product alerts (UTF-8).""" | |
| raw = json.dumps(alert.to_dict(), separators=(",", ":"), ensure_ascii=True).encode("utf-8") | |
| if len(raw) > MAX_ALERT_BYTES: | |
| raise ValueError(f"alert payload exceeds MAX_ALERT_BYTES ({len(raw)})") | |
| return raw | |
| def bytes_to_alert(data: bytes) -> ProductAlert: | |
| """Parse alert from wire or inner JSON bytes.""" | |
| try: | |
| inner = unpack_wire(data) | |
| except ValueError: | |
| inner = data | |
| # If unpack produced garbage for pure JSON, try raw | |
| try: | |
| obj = json.loads(inner.decode("utf-8")) | |
| except (UnicodeDecodeError, json.JSONDecodeError): | |
| obj = json.loads(data.decode("utf-8")) | |
| if not isinstance(obj, dict): | |
| raise ValueError("alert JSON must be an object") | |
| required = ("zone_id", "product_actionable", "elevated", "alert_level") | |
| for k in required: | |
| if k not in obj: | |
| raise ValueError(f"alert missing required field: {k}") | |
| return ProductAlert( | |
| zone_id=str(obj["zone_id"]), | |
| product_actionable=bool(obj["product_actionable"]), | |
| elevated=bool(obj["elevated"]), | |
| alert_level=str(obj["alert_level"]), | |
| drought_risk=float(obj.get("drought_risk", 0.0)), | |
| flood_risk=float(obj.get("flood_risk", 0.0)), | |
| max_risk=float(obj.get("max_risk", 0.0)), | |
| confidence=float(obj.get("confidence", 0.0)), | |
| trigger=str(obj.get("trigger", "")), | |
| ts=float(obj.get("ts", time.time())), | |
| schema_version=int(obj.get("schema_version", 1)), | |
| extras=dict(obj.get("extras") or {}), | |
| ) | |
| def product_alert_from_risk( | |
| zone_id: str, | |
| *, | |
| alert_level: str, | |
| product_actionable: bool, | |
| elevated: bool, | |
| drought_risk: float = 0.0, | |
| flood_risk: float = 0.0, | |
| max_risk: float = 0.0, | |
| confidence: float = 0.0, | |
| trigger: str = "", | |
| **extra: Any, | |
| ) -> ProductAlert: | |
| """Helper for env / edge callers building an alert from scorer/env info.""" | |
| return ProductAlert( | |
| zone_id=zone_id, | |
| product_actionable=product_actionable, | |
| elevated=elevated, | |
| alert_level=alert_level, | |
| drought_risk=drought_risk, | |
| flood_risk=flood_risk, | |
| max_risk=max_risk, | |
| confidence=confidence, | |
| trigger=trigger, | |
| extras=dict(extra) if extra else {}, | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Protocol | |
| # --------------------------------------------------------------------------- | |
| class NodeTransport(Protocol): | |
| """Minimal interface every transport must implement. | |
| Hidden state and product alerts are peer channels with identical | |
| latest-value / TTL semantics. Topic names, brokers, and P2P routes | |
| must not leak into callers. | |
| """ | |
| async def send_hidden_state(self, zone_id: str, hidden: bytes) -> None: | |
| ... | |
| async def recv_hidden_state(self, zone_id: str) -> Optional[bytes]: | |
| ... | |
| async def send_alert(self, zone_id: str, alert: bytes) -> None: | |
| ... | |
| async def recv_alert(self, zone_id: str) -> Optional[bytes]: | |
| ... | |
| async def clear(self) -> None: | |
| ... | |
| # --------------------------------------------------------------------------- | |
| # LocalTransport — training + testing | |
| # --------------------------------------------------------------------------- | |
| class LocalTransport: | |
| """ | |
| In-process deterministic transport. | |
| Semantics: | |
| - Latest-value (NOT consuming) | |
| - TTL enforced on recv | |
| - Auto wire-envelope on send when payload is bare inner | |
| """ | |
| def __init__(self, *, compress: bool = False, ttl_sec: float = STATE_TTL_SEC) -> None: | |
| self._hidden: Dict[str, Tuple[bytes, float]] = {} | |
| self._alerts: Dict[str, Tuple[bytes, float]] = {} | |
| self._compress = compress | |
| self._ttl = float(ttl_sec) | |
| def _pack(self, payload: bytes) -> bytes: | |
| # If already versioned, leave as-is | |
| if len(payload) >= 2 and payload[0] == WIRE_VERSION and not _looks_like_legacy_hidden(payload): | |
| return payload | |
| return pack_wire(payload, compress=self._compress) | |
| async def send_hidden_state(self, zone_id: str, hidden: bytes) -> None: | |
| if len(hidden) > MAX_HIDDEN_BYTES + 16: | |
| raise ValueError("Hidden state exceeds MAX_HIDDEN_BYTES") | |
| wire = self._pack(hidden) | |
| self._hidden[zone_id] = (wire, time.time()) | |
| logger.debug("LocalTransport: stored hidden %d bytes for %s", len(wire), zone_id) | |
| async def recv_hidden_state(self, zone_id: str) -> Optional[bytes]: | |
| item = self._hidden.get(zone_id) | |
| if item is None: | |
| return None | |
| data, ts = item | |
| if time.time() - ts > self._ttl: | |
| logger.debug("LocalTransport: hidden expired for %s", zone_id) | |
| return None | |
| return data | |
| async def send_alert(self, zone_id: str, alert: bytes) -> None: | |
| if len(alert) > MAX_ALERT_BYTES + 16: | |
| raise ValueError("Alert exceeds MAX_ALERT_BYTES") | |
| wire = self._pack(alert) | |
| self._alerts[zone_id] = (wire, time.time()) | |
| logger.debug("LocalTransport: stored alert %d bytes for %s", len(wire), zone_id) | |
| async def recv_alert(self, zone_id: str) -> Optional[bytes]: | |
| item = self._alerts.get(zone_id) | |
| if item is None: | |
| return None | |
| data, ts = item | |
| if time.time() - ts > self._ttl: | |
| logger.debug("LocalTransport: alert expired for %s", zone_id) | |
| return None | |
| return data | |
| async def clear(self) -> None: | |
| self._hidden.clear() | |
| self._alerts.clear() | |
| # --------------------------------------------------------------------------- | |
| # MQTTTransport — edge fleet | |
| # --------------------------------------------------------------------------- | |
| class MQTTTransport: | |
| """ | |
| MQTT backend for NodeTransport (hub-and-spoke). | |
| Guarantees: | |
| - Same latest-value / non-consuming / TTL semantics as LocalTransport | |
| - Thread-safe receive buffers | |
| - Fallback to LocalTransport when broker unavailable | |
| - TLS + username/password when configured | |
| - LWT + retained birth on presence topic | |
| - MQTT 5 message expiry when broker/protocol supports it; 3.1.1 fallback | |
| - Multi-broker: try hosts in order (step 5 topology as config) | |
| Topic layout (implementation detail — not part of NodeTransport): | |
| {topic_root}/hidden/{zone_id} | |
| {topic_root}/alert/{zone_id} | |
| {topic_root}/presence/{node_id} | |
| """ | |
| def __init__( | |
| self, | |
| broker: str = "localhost", | |
| port: int = 1883, | |
| username: Optional[str] = None, | |
| password: Optional[str] = None, | |
| topic_root: str = DEFAULT_TOPIC_ROOT, | |
| # legacy alias | |
| topic_prefix: Optional[str] = None, | |
| *, | |
| # step 5: try multiple brokers in order | |
| brokers: Optional[Sequence[Tuple[str, int]]] = None, | |
| # step 2: TLS | |
| tls: bool = False, | |
| tls_ca_certs: Optional[str] = None, | |
| tls_certfile: Optional[str] = None, | |
| tls_keyfile: Optional[str] = None, | |
| tls_insecure: bool = False, | |
| # identity / presence | |
| client_id: Optional[str] = None, | |
| node_id: str = "node-0", | |
| # behaviour | |
| compress: bool = False, | |
| ttl_sec: float = STATE_TTL_SEC, | |
| keepalive: int = 60, | |
| use_mqttv5: bool = True, | |
| qos: int = 1, | |
| ) -> None: | |
| if brokers: | |
| self._brokers: List[Tuple[str, int]] = [(str(h), int(p)) for h, p in brokers] | |
| else: | |
| self._brokers = [(broker, port)] | |
| self.username = username | |
| self.password = password | |
| self.topic_root = topic_root.rstrip("/") | |
| # backward compat: old topic_prefix meant hidden prefix | |
| if topic_prefix is not None: | |
| # if caller passed weather/hidden, derive root | |
| parts = topic_prefix.rstrip("/").split("/") | |
| if parts and parts[-1] == "hidden": | |
| self.topic_root = "/".join(parts[:-1]) or DEFAULT_TOPIC_ROOT | |
| else: | |
| self.topic_root = topic_prefix.rstrip("/") | |
| self.tls = tls | |
| self.tls_ca_certs = tls_ca_certs | |
| self.tls_certfile = tls_certfile | |
| self.tls_keyfile = tls_keyfile | |
| self.tls_insecure = tls_insecure | |
| self.client_id = client_id | |
| self.node_id = node_id | |
| self.compress = compress | |
| self._ttl = float(ttl_sec) | |
| self.keepalive = keepalive | |
| self.use_mqttv5 = use_mqttv5 | |
| self.qos = int(qos) | |
| self._client = None | |
| self._connected = False | |
| self._mqttv5_active = False | |
| self._active_broker: Optional[Tuple[str, int]] = None | |
| self._hidden_rx: Dict[str, Tuple[bytes, float]] = {} | |
| self._alert_rx: Dict[str, Tuple[bytes, float]] = {} | |
| self._lock = threading.Lock() | |
| self._fallback = LocalTransport(compress=compress, ttl_sec=ttl_sec) | |
| # -- topic helpers (MQTT-only) ------------------------------------------ | |
| def _hidden_topic(self, zone_id: str) -> str: | |
| return f"{self.topic_root}/hidden/{zone_id}" | |
| def _alert_topic(self, zone_id: str) -> str: | |
| return f"{self.topic_root}/alert/{zone_id}" | |
| def _presence_topic(self) -> str: | |
| return f"{self.topic_root}/presence/{self.node_id}" | |
| def _pack(self, payload: bytes) -> bytes: | |
| if len(payload) >= 2 and payload[0] == WIRE_VERSION and not _looks_like_legacy_hidden(payload): | |
| return payload | |
| return pack_wire(payload, compress=self.compress) | |
| # -- connect / disconnect ----------------------------------------------- | |
| async def connect(self) -> bool: | |
| try: | |
| import paho.mqtt.client as mqtt | |
| except ImportError: | |
| logger.warning("paho-mqtt not installed — using LocalTransport fallback only") | |
| return False | |
| last_err: Optional[Exception] = None | |
| for host, port in self._brokers: | |
| try: | |
| ok = self._connect_one(mqtt, host, port) | |
| if ok: | |
| self._active_broker = (host, port) | |
| return True | |
| except Exception as e: | |
| last_err = e | |
| logger.warning("MQTTTransport: broker %s:%s failed: %s", host, port, e) | |
| self._teardown_client() | |
| logger.error("MQTTTransport: all brokers failed (last=%s)", last_err) | |
| return False | |
| def _connect_one(self, mqtt, host: str, port: int) -> bool: | |
| protocol = mqtt.MQTTv311 | |
| callback_api = getattr(mqtt, "CallbackAPIVersion", None) | |
| client_kwargs: Dict[str, Any] = {} | |
| if self.client_id: | |
| client_kwargs["client_id"] = self.client_id | |
| # Prefer MQTT 5 when requested and available | |
| self._mqttv5_active = False | |
| if self.use_mqttv5 and hasattr(mqtt, "MQTTv5"): | |
| protocol = mqtt.MQTTv5 | |
| self._mqttv5_active = True | |
| client_kwargs["protocol"] = protocol | |
| if callback_api is not None: | |
| # paho 2.x | |
| try: | |
| client = mqtt.Client( | |
| callback_api_version=callback_api.VERSION2, | |
| **client_kwargs, | |
| ) | |
| except TypeError: | |
| client = mqtt.Client(**client_kwargs) | |
| else: | |
| client = mqtt.Client(**client_kwargs) | |
| if self.username is not None: | |
| client.username_pw_set(self.username, self.password) | |
| if self.tls: | |
| tls_args: Dict[str, Any] = {} | |
| if self.tls_ca_certs: | |
| tls_args["ca_certs"] = self.tls_ca_certs | |
| if self.tls_certfile: | |
| tls_args["certfile"] = self.tls_certfile | |
| if self.tls_keyfile: | |
| tls_args["keyfile"] = self.tls_keyfile | |
| client.tls_set(**tls_args) if tls_args else client.tls_set() | |
| if self.tls_insecure: | |
| client.tls_insecure_set(True) | |
| # LWT: retained offline presence | |
| lwt_payload = json.dumps( | |
| {"node_id": self.node_id, "status": "offline", "ts": time.time()}, | |
| separators=(",", ":"), | |
| ).encode("utf-8") | |
| client.will_set( | |
| self._presence_topic(), | |
| payload=lwt_payload, | |
| qos=self.qos, | |
| retain=True, | |
| ) | |
| client.on_connect = self._on_connect | |
| client.on_disconnect = self._on_disconnect | |
| client.on_message = self._on_message | |
| # Automatic reconnect (paho) | |
| try: | |
| client.reconnect_delay_set(min_delay=1, max_delay=30) | |
| except Exception: | |
| pass | |
| client.connect(host, port, keepalive=self.keepalive) | |
| client.subscribe(f"{self.topic_root}/hidden/#", qos=self.qos) | |
| client.subscribe(f"{self.topic_root}/alert/#", qos=self.qos) | |
| client.loop_start() | |
| self._client = client | |
| # Brief wait for on_connect — non-blocking best-effort | |
| for _ in range(20): | |
| if self._connected: | |
| break | |
| time.sleep(0.05) | |
| logger.info( | |
| "MQTTTransport: connect attempted %s:%s mqttv5=%s connected=%s", | |
| host, port, self._mqttv5_active, self._connected, | |
| ) | |
| # Consider success if loop started; on_connect may race | |
| return True | |
| def _teardown_client(self) -> None: | |
| if self._client is not None: | |
| try: | |
| self._client.loop_stop() | |
| self._client.disconnect() | |
| except Exception: | |
| pass | |
| self._client = None | |
| self._connected = False | |
| def _on_connect(self, client, userdata, flags, reason_code, properties=None) -> None: | |
| # Compatible with both paho 1.x (rc int) and 2.x VERSION2 (reason_code) | |
| rc = reason_code | |
| if hasattr(reason_code, "value"): | |
| rc = reason_code.value | |
| elif hasattr(reason_code, "is_failure"): | |
| rc = 0 if not reason_code.is_failure else 1 | |
| self._connected = (rc == 0) if isinstance(rc, int) else bool(rc == 0) | |
| logger.info("MQTTTransport: on_connect rc=%s connected=%s", rc, self._connected) | |
| if not self._connected: | |
| return | |
| # Retained birth / online presence | |
| birth = json.dumps( | |
| {"node_id": self.node_id, "status": "online", "ts": time.time()}, | |
| separators=(",", ":"), | |
| ).encode("utf-8") | |
| try: | |
| client.publish( | |
| self._presence_topic(), | |
| birth, | |
| qos=self.qos, | |
| retain=True, | |
| ) | |
| except Exception as e: | |
| logger.warning("MQTTTransport: birth publish failed: %s", e) | |
| # Re-subscribe after reconnect | |
| try: | |
| client.subscribe(f"{self.topic_root}/hidden/#", qos=self.qos) | |
| client.subscribe(f"{self.topic_root}/alert/#", qos=self.qos) | |
| except Exception as e: | |
| logger.warning("MQTTTransport: resubscribe failed: %s", e) | |
| def _on_disconnect(self, client, userdata, flags, reason_code=None, properties=None) -> None: | |
| # paho 1.x: (client, userdata, rc); 2.x VERSION2 adds flags/properties | |
| self._connected = False | |
| logger.warning("MQTTTransport: disconnected reason=%s", reason_code) | |
| def _on_message(self, client, userdata, msg) -> None: | |
| try: | |
| payload = bytes(msg.payload) | |
| if len(payload) > MAX_HIDDEN_BYTES + 64: | |
| logger.warning("MQTTTransport: payload too large — dropped") | |
| return | |
| topic = msg.topic or "" | |
| parts = topic.split("/") | |
| if len(parts) < 3: | |
| return | |
| kind = parts[-2] # hidden | alert | |
| zone_id = parts[-1] | |
| if not zone_id or zone_id == self.node_id and kind == "presence": | |
| return | |
| now = time.time() | |
| with self._lock: | |
| if kind == "hidden": | |
| self._hidden_rx[zone_id] = (payload, now) | |
| elif kind == "alert": | |
| self._alert_rx[zone_id] = (payload, now) | |
| else: | |
| return | |
| logger.debug( | |
| "MQTTTransport: received %s %d bytes for %s", | |
| kind, len(payload), zone_id, | |
| ) | |
| except Exception as e: | |
| logger.warning("MQTTTransport._on_message error: %s", e) | |
| def _publish(self, topic: str, payload: bytes) -> None: | |
| assert self._client is not None | |
| kwargs: Dict[str, Any] = { | |
| "topic": topic, | |
| "payload": payload, | |
| "qos": self.qos, | |
| "retain": True, | |
| } | |
| if self._mqttv5_active: | |
| try: | |
| from paho.mqtt.properties import Properties | |
| from paho.mqtt.packettypes import PacketTypes | |
| props = Properties(PacketTypes.PUBLISH) | |
| # Message expiry (seconds) — MQTT 5 equivalent of soft TTL | |
| props.MessageExpiryInterval = int(self._ttl) | |
| # User properties for schema / debugging (not required by receivers) | |
| try: | |
| props.UserProperty = [("wire_v", str(WIRE_VERSION))] | |
| except Exception: | |
| pass | |
| kwargs["properties"] = props | |
| except Exception as e: | |
| logger.debug("MQTTTransport: MQTT5 properties unavailable: %s", e) | |
| self._client.publish(**kwargs) | |
| # -- NodeTransport API -------------------------------------------------- | |
| async def send_hidden_state(self, zone_id: str, hidden: bytes) -> None: | |
| if len(hidden) > MAX_HIDDEN_BYTES + 16: | |
| raise ValueError("Hidden state exceeds MAX_HIDDEN_BYTES") | |
| wire = self._pack(hidden) | |
| if self._client is None or not self._connected: | |
| await self._fallback.send_hidden_state(zone_id, wire) | |
| return | |
| try: | |
| self._publish(self._hidden_topic(zone_id), wire) | |
| logger.debug("MQTTTransport: published hidden to %s", zone_id) | |
| except Exception as e: | |
| logger.warning("MQTTTransport: hidden publish failed (%s) — fallback", e) | |
| await self._fallback.send_hidden_state(zone_id, wire) | |
| async def recv_hidden_state(self, zone_id: str) -> Optional[bytes]: | |
| if self._client is None or not self._connected: | |
| return await self._fallback.recv_hidden_state(zone_id) | |
| with self._lock: | |
| item = self._hidden_rx.get(zone_id) | |
| if item is None: | |
| # also check fallback (local writes while disconnected) | |
| return await self._fallback.recv_hidden_state(zone_id) | |
| data, ts = item | |
| if time.time() - ts > self._ttl: | |
| logger.debug("MQTTTransport: hidden expired for %s", zone_id) | |
| return None | |
| return data | |
| async def send_alert(self, zone_id: str, alert: bytes) -> None: | |
| if len(alert) > MAX_ALERT_BYTES + 16: | |
| raise ValueError("Alert exceeds MAX_ALERT_BYTES") | |
| wire = self._pack(alert) | |
| if self._client is None or not self._connected: | |
| await self._fallback.send_alert(zone_id, wire) | |
| return | |
| try: | |
| self._publish(self._alert_topic(zone_id), wire) | |
| logger.debug("MQTTTransport: published alert to %s", zone_id) | |
| except Exception as e: | |
| logger.warning("MQTTTransport: alert publish failed (%s) — fallback", e) | |
| await self._fallback.send_alert(zone_id, wire) | |
| async def recv_alert(self, zone_id: str) -> Optional[bytes]: | |
| if self._client is None or not self._connected: | |
| return await self._fallback.recv_alert(zone_id) | |
| with self._lock: | |
| item = self._alert_rx.get(zone_id) | |
| if item is None: | |
| return await self._fallback.recv_alert(zone_id) | |
| data, ts = item | |
| if time.time() - ts > self._ttl: | |
| logger.debug("MQTTTransport: alert expired for %s", zone_id) | |
| return None | |
| return data | |
| async def clear(self) -> None: | |
| with self._lock: | |
| self._hidden_rx.clear() | |
| self._alert_rx.clear() | |
| await self._fallback.clear() | |
| async def disconnect(self) -> None: | |
| if self._client is not None: | |
| # Best-effort offline presence before disconnect | |
| try: | |
| offline = json.dumps( | |
| {"node_id": self.node_id, "status": "offline", "ts": time.time()}, | |
| separators=(",", ":"), | |
| ).encode("utf-8") | |
| self._client.publish( | |
| self._presence_topic(), | |
| offline, | |
| qos=self.qos, | |
| retain=True, | |
| ) | |
| except Exception: | |
| pass | |
| self._teardown_client() | |
| self._active_broker = None | |
| # --------------------------------------------------------------------------- | |
| # Factory | |
| # --------------------------------------------------------------------------- | |
| def create_node_transport( | |
| use_mqtt: bool = False, | |
| **mqtt_kwargs: Any, | |
| ) -> NodeTransport: | |
| """ | |
| Public factory. | |
| Usage: | |
| training → create_node_transport() | |
| edge → create_node_transport(use_mqtt=True, broker="...", tls=True, ...) | |
| multi → create_node_transport(use_mqtt=True, brokers=[("a", 8883), ("b", 8883)], tls=True) | |
| """ | |
| if use_mqtt: | |
| return MQTTTransport(**mqtt_kwargs) | |
| return LocalTransport( | |
| compress=bool(mqtt_kwargs.get("compress", False)), | |
| ttl_sec=float(mqtt_kwargs.get("ttl_sec", STATE_TTL_SEC)), | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Offline self-test | |
| # --------------------------------------------------------------------------- | |
| def _self_test() -> None: | |
| import asyncio | |
| async def _run() -> None: | |
| print("node_transport.py self-test") | |
| # --- wire pack/unpack --- | |
| raw = b"hello-alert-payload" | |
| wire = pack_wire(raw, compress=False) | |
| assert unpack_wire(wire) == raw | |
| wire_c = pack_wire(raw * 20, compress=True) | |
| assert unpack_wire(wire_c) == raw * 20 | |
| print(" wire envelope OK") | |
| # --- hidden round-trip (no torch required if we skip tensor path) --- | |
| # Build a synthetic inner payload | |
| n_layers, batch, hidden = 1, 1, 8 | |
| inner = struct.pack("<III", n_layers, batch, hidden) + (b"\x00\x00\x00\x00" * (n_layers * batch * hidden)) | |
| assert _looks_like_legacy_hidden(inner) | |
| w = pack_wire(inner, compress=True) | |
| assert unpack_wire(w) == inner | |
| print(" hidden wire OK") | |
| # --- alert serialisation --- | |
| alert = product_alert_from_risk( | |
| "karawang_rice", | |
| alert_level="warning", | |
| product_actionable=True, | |
| elevated=True, | |
| drought_risk=0.42, | |
| flood_risk=0.05, | |
| trigger="drought_warning", | |
| confidence=0.63, | |
| ) | |
| ab = alert_to_bytes(alert) | |
| alert2 = bytes_to_alert(pack_wire(ab)) | |
| assert alert2.zone_id == "karawang_rice" | |
| assert alert2.product_actionable is True | |
| assert alert2.alert_level == "warning" | |
| assert abs(alert2.drought_risk - 0.42) < 1e-6 | |
| print(" alert serialisation OK") | |
| # --- LocalTransport protocol --- | |
| tr = LocalTransport() | |
| await tr.send_hidden_state("z1", inner) | |
| got = await tr.recv_hidden_state("z1") | |
| assert got is not None | |
| assert unpack_wire(got) == inner | |
| await tr.send_alert("z1", ab) | |
| got_a = await tr.recv_alert("z1") | |
| assert got_a is not None | |
| a3 = bytes_to_alert(got_a) | |
| assert a3.product_actionable and a3.zone_id == "karawang_rice" | |
| await tr.clear() | |
| assert await tr.recv_hidden_state("z1") is None | |
| assert await tr.recv_alert("z1") is None | |
| print(" LocalTransport protocol OK") | |
| # --- MQTTTransport falls back without broker --- | |
| mqtt_tr = MQTTTransport(broker="127.0.0.1", port=1, node_id="test-node") | |
| # do not connect — should use fallback | |
| await mqtt_tr.send_alert("z2", ab) | |
| got_b = await mqtt_tr.recv_alert("z2") | |
| assert got_b is not None | |
| assert bytes_to_alert(got_b).alert_level == "warning" | |
| await mqtt_tr.send_hidden_state("z2", inner) | |
| assert await mqtt_tr.recv_hidden_state("z2") is not None | |
| print(" MQTTTransport offline fallback OK") | |
| # --- factory --- | |
| t0 = create_node_transport() | |
| assert isinstance(t0, LocalTransport) | |
| t1 = create_node_transport(use_mqtt=True, broker="localhost") | |
| assert isinstance(t1, MQTTTransport) | |
| print(" factory OK") | |
| print("All node_transport self-tests passed.") | |
| asyncio.run(_run()) | |
| if __name__ == "__main__": | |
| logging.basicConfig(level=logging.INFO) | |
| _self_test() | |