"""MQTT nodes: a subscriber that wakes the graph, a publisher that speaks for it.""" from __future__ import annotations import asyncio import json import logging from collections.abc import Iterable from enum import Enum from typing import TYPE_CHECKING, Any from pydantic import BaseModel, ConfigDict, Field from fluksio.flow.messages import MessageSpec from fluksio.flow.nodes.base import Node if TYPE_CHECKING: from fastapi import FastAPI logger = logging.getLogger(__name__) def topic_matches(filter_: str, topic: str) -> bool: """Does an MQTT topic filter cover this topic? The broker already decided to send it, so this is only about which port it belongs to. ``+`` stands for one level, ``#`` for the rest of them; a filter with neither is an exact string, which is the common case and the one the caller checks first. """ if filter_ == topic: return True parts = filter_.split("/") levels = topic.split("/") for i, part in enumerate(parts): if part == "#": # Everything before it already matched, and '#' takes the rest — # including nothing at all, so 'a/#' covers 'a' as the spec says. return True if i >= len(levels): return False if part != "+" and part != levels[i]: return False return len(parts) == len(levels) class MqttNode(Node): """ MQTT node that can act as a subscriber (trigger) or publisher (sender). This node integrates with an MQTT broker to either: - **Trigger mode (Subscriber)**: Subscribe to MQTT topics and inject received messages into the pipeline. Used when ``provides`` is specified but ``requires`` is empty. - **Sender mode (Publisher)**: Publish pipeline data to MQTT topics. Used when ``requires`` is specified. The ``topic`` parameter in ``params`` controls the mapping between pipeline message names and MQTT topics: - **dict**: Explicit mapping from message name to MQTT topic, e.g. ``{"temperature": "sensors/room1/temp", "humidity": "sensors/room1/hum"}``. - **str** (legacy): A single topic string. All messages are mapped to this one topic (subscriber receives from it, publisher sends to it). :param requires: Messages required by this node (makes it a publisher node). :type requires: MessageSpec | list[MessageSpec] | None :param provides: Messages provided by this node (makes it a subscriber node). :type provides: MessageSpec | list[MessageSpec] | None :param params: Parameters dict containing: - ``topic`` (str | dict): MQTT topic(s). A dict maps message names to individual topics. A plain string uses that topic for all messages. - ``broker_host`` (str): MQTT broker hostname (default: "localhost") - ``broker_port`` (int): MQTT broker port (default: 1883) - ``username`` (str | None): Optional username for authentication - ``password`` (str | None): Optional password for authentication - ``client_id`` (str | None): Optional client ID - ``qos`` (int): Quality of Service level 0, 1, or 2 (default: 0) - ``retain`` (bool): Retain flag for published messages (default: False) - ``keepalive`` (int): Keepalive interval in seconds (default: 60) - ``timeout`` (float): Deadline for a broker operation in seconds (default: 10.0) - ``publish_queue_size`` (int): Publisher backlog depth (default: 256) :type params: dict :param name: Optional name for the node. :type name: str | None :raises ValueError: If both ``requires`` and ``provides`` are empty. :example: Subscriber with per-message topics: >>> subscriber = MqttNode( ... provides=[ ... MessageSpec(name="inverter_input", dtype=DType.FLOAT), ... MessageSpec(name="inverter_output", dtype=DType.FLOAT), ... ], ... params={ ... "topic": { ... "inverter_input": "sensors/pv", ... "inverter_output": "sensors/output", ... }, ... "broker_host": "localhost", ... }, ... ) Publisher with per-message topics: >>> publisher = MqttNode( ... requires=[ ... MessageSpec(name="target_temp", dtype=DType.FLOAT), ... MessageSpec(name="fan_speed", dtype=DType.INT), ... ], ... params={ ... "topic": { ... "target_temp": "actuators/hvac/temp", ... "fan_speed": "actuators/hvac/fan", ... }, ... "broker_host": "localhost", ... "qos": 1, ... }, ... ) Legacy single-topic subscriber: >>> subscriber = MqttNode( ... provides=[MessageSpec(name="temperature", dtype=DType.FLOAT)], ... params={"topic": "sensors/temperature", "broker_host": "localhost"}, ... ) """ class Mode(Enum): """Operating mode of the MQTT node.""" SUBSCRIBER = "subscriber" # Receives MQTT messages (trigger) PUBLISHER = "publisher" # Sends MQTT messages (sender) # Publishing again is a second command to whatever is listening. idempotent = False __slots__ = ( "topics", "mode", "broker_host", "broker_port", "username", "password", "client_id", "qos", "retain", "keepalive", "timeout", "publish_queue_size", "json_keys", "_topic_to_ports", "_wildcards", "_subscription_task", "_mqtt_client", "_stop_event", "_publish_queue", "_publisher_task", "_loop", ) class Params(BaseModel): model_config = ConfigDict(extra="allow") # One topic for every port, or a per-port mapping. topic: str | dict[str, str] = "*" broker_host: str = "localhost" broker_port: int = 1883 username: str | None = None password: str | None = Field(default=None, json_schema_extra={"x-secret": True}) client_id: str | None = None qos: int = 0 retain: bool = False keepalive: int = 60 # Bounds every broker operation: subscribe, publish, and the # disconnect acknowledgement on the way out. Without one a client # whose socket died waits for that ack forever, and the task never # finishes unwinding. Brokers differ, so it is per node. timeout: float = Field( default=10.0, gt=0, description="Give up on a broker operation after this many seconds.", ) # Deep enough to ride out a broker hiccup, shallow enough that a # publisher which cannot keep up drops old values instead of growing # without bound. A node that bursts wants more than one that trickles. publish_queue_size: int = Field( default=256, gt=0, description=( "How many payloads may wait for the broker. Past this the oldest " "is dropped and the node reports degraded." ), ) # Which key to lift out of a JSON object payload. A device that wraps # its reading — Victron's ``{"value": 5}`` — is otherwise a Python node # per port. One key for every port, or a per-port mapping. json_key: str | dict[str, str] = "" @classmethod def instance_key(cls, params: dict[str, Any]) -> str | None: """The broker and topic, which is one physical thing. A publisher and a subscriber on the same topic get the same key on purpose: they are two ends of one wire, and drawing them as one neuron is the only way the path through the broker shows up at all. """ # ponytail: publisher and subscriber merge into one neuron; key on mode # as well if the two directions ever need telling apart. fields = cls.Params.model_fields topic = params.get("topic", fields["topic"].default) if isinstance(topic, dict): topic = json.dumps(topic, sort_keys=True) host = params.get("broker_host", fields["broker_host"].default) port = params.get("broker_port", fields["broker_port"].default) return f"{host}:{port}/{topic}" def __init__( self, requires: MessageSpec | Iterable[MessageSpec] = (), provides: MessageSpec | Iterable[MessageSpec] = (), params: dict[str, Any] | None = None, name: str | None = None, ): cfg = self.Params.model_validate(params or {}) requires = Node._normalize_ports(requires) provides = Node._normalize_ports(provides) if not requires and not provides: raise ValueError( "An MQTT node needs either inputs (to publish) or outputs " "(to subscribe)" ) # Inputs mean this node publishes; outputs mean it subscribes. self.mode = MqttNode.Mode.PUBLISHER if requires else MqttNode.Mode.SUBSCRIBER ports = provides if self.mode == MqttNode.Mode.SUBSCRIBER else requires if isinstance(cfg.topic, dict): self.topics: dict[str, str] = dict(cfg.topic) else: self.topics = {spec.port: cfg.topic for spec in ports} # Reverse lookup for routing incoming payloads back to ports. A filter # holding a wildcard cannot be found by the incoming topic, so those # are kept aside and walked when the exact lookup misses. self._topic_to_ports: dict[str, list[str]] = {} for port, topic in self.topics.items(): self._topic_to_ports.setdefault(topic, []).append(port) self._wildcards = [t for t in self._topic_to_ports if "+" in t or "#" in t] if isinstance(cfg.json_key, dict): self.json_keys: dict[str, str] = dict(cfg.json_key) else: self.json_keys = ( {spec.port: cfg.json_key for spec in ports} if cfg.json_key else {} ) self.broker_host = cfg.broker_host self.broker_port = cfg.broker_port self.username = cfg.username self.password = cfg.password self.client_id = cfg.client_id self.qos = cfg.qos self.retain = cfg.retain self.keepalive = cfg.keepalive self.timeout = cfg.timeout self.publish_queue_size = cfg.publish_queue_size # Runtime state self._subscription_task: asyncio.Task[None] | None = None self._mqtt_client = None self._stop_event: asyncio.Event | None = None self._publish_queue: asyncio.Queue[dict[str, Any]] | None = None self._publisher_task: asyncio.Task[None] | None = None self._loop: asyncio.AbstractEventLoop | None = None # Set default name based on mode and topics if name is None: unique_topics = set(self.topics.values()) if len(unique_topics) == 1: safe_topic = ( next(iter(unique_topics)) .replace("/", "_") .replace("+", "x") .replace("#", "all") .strip("_") ) else: safe_topic = f"{len(unique_topics)}topics" name = f"mqtt_{self.mode.value}_{safe_topic}" # Initialize parent with appropriate function # For subscriber mode, f is a no-op since data is injected via inject() # For publisher mode, f handles the outgoing MQTT publish super().__init__( f=( self._noop_subscriber if self.mode == MqttNode.Mode.SUBSCRIBER else self._publisher_handler ), requires=requires, provides=provides, params=params, name=name, ) @staticmethod def _noop_subscriber( params: dict[str, Any], **kwargs: Any ) -> dict[str, Any] | None: """ No-op function for subscriber mode nodes. Subscriber mode nodes inject data via :meth:`inject`, not :meth:`__call__`. This function exists only to satisfy the Node interface and should not be called directly. :param params: Node parameters (unused). :type params: dict :param kwargs: Additional arguments (unused). :type kwargs: Any :returns: Always returns None. :rtype: None """ return None def _publisher_handler( self, params: dict[str, Any], **kwargs: Any ) -> dict[str, Any] | None: """ Publish pipeline data to MQTT topic (publisher mode). This method is called when upstream dependencies are satisfied. Handing the payload to the node's publisher task is all that happens here: the task holds one connection for the node's lifetime, where connecting per message would cost a full handshake every time. :param params: Node parameters. :type params: dict :param kwargs: Pipeline data to publish (from required messages). :type kwargs: Any :returns: None (publishing is fire-and-forget). :rtype: dict | None """ loop, queue = self._loop, self._publish_queue if loop is not None and queue is not None: loop.call_soon_threadsafe(self._enqueue, queue, dict(kwargs)) return None # No publisher task: a node built for a draft preview or a test. Send it # the one-shot way rather than silently dropping the message. try: asyncio.run(self._publish_once(kwargs)) except RuntimeError: logger.warning( "Node '%s' cannot publish from a running event loop unstarted", self.name, ) return None def _enqueue( self, queue: asyncio.Queue[dict[str, Any]], data: dict[str, Any] ) -> None: """Queue a payload, dropping the oldest when the broker cannot keep up.""" if queue.full(): try: queue.get_nowait() logger.warning("Publish queue full for node '%s', dropped", self.name) self.report_health("degraded", "publish queue full") except asyncio.QueueEmpty: pass queue.put_nowait(data) async def _publisher_loop(self) -> None: """Hold one connection and drain the publish queue over it. A dropped connection raises, and the supervisor decides when to reconnect — the same arrangement the subscriber uses. """ import aiomqtt queue = self._publish_queue if queue is None: return async with aiomqtt.Client( hostname=self.broker_host, port=self.broker_port, username=self.username, password=self.password, identifier=self.client_id, keepalive=self.keepalive, timeout=self.timeout, ) as client: self.report_health("ok") while True: data = await queue.get() try: await self._publish_with(client, data) except Exception as exc: self.report_health("down", str(exc)) raise async def _publish_once(self, data: dict[str, Any]) -> None: """Connect, publish, disconnect — the unstarted node's path.""" import aiomqtt async with aiomqtt.Client( hostname=self.broker_host, port=self.broker_port, username=self.username, password=self.password, identifier=self.client_id, keepalive=self.keepalive, timeout=self.timeout, ) as client: await self._publish_with(client, data) async def _publish_with(self, client: Any, data: dict[str, Any]) -> None: """ Publish messages to their mapped MQTT topics. Each message in *data* is published to its corresponding topic from the ``topics`` mapping. Messages are sent as individual JSON payloads per topic. :param data: Data to publish, keyed by port name. :type data: dict """ import json for port, value in data.items(): topic = self.topics.get(port) if topic is None: logger.warning( "No topic mapping for port '%s' in node '%s', skipping", port, self.name, ) continue # A string goes on the wire as it stands. Devices on a shared # broker expect bare values, and the subscriber below already # falls back to the raw text when it is not JSON, so a # fluksio-to-fluksio round trip is unaffected. payload = value if isinstance(value, str) else json.dumps(value) await client.publish( topic, payload=payload, qos=self.qos, retain=self.retain, ) logger.info( "Published to '%s' from node '%s': %s", topic, self.name, payload, ) async def start(self, app: FastAPI | None = None) -> None: """A subscriber listens; a publisher opens the connection it will reuse.""" if self.mode is MqttNode.Mode.SUBSCRIBER: await self.start_subscription() else: await self.start_publisher() async def stop(self, app: FastAPI | None = None) -> None: await self.stop_subscription() await self.stop_publisher() async def start_publisher(self) -> None: """Run the task that owns this node's connection to the broker.""" if self._publish_queue is not None: return self._publish_queue = asyncio.Queue(maxsize=self.publish_queue_size) self._loop = asyncio.get_running_loop() self._publisher_task = self._run_supervised("mqtt-out", self._publisher_loop) async def stop_publisher(self) -> None: """Drop the queue and let the connection go.""" if self._publish_queue is None: return if self._publisher_task is not None: await self._cancel_task(self._publisher_task) self._publisher_task = None self._publish_queue = None self._loop = None async def start_subscription(self) -> None: """ Start the MQTT subscription for trigger mode nodes. This method starts a background task that listens for messages on the subscribed topic and triggers the pipeline when messages arrive. :raises RuntimeError: If called on a publisher mode node. :example: >>> subscriber = MqttNode( ... topic="sensors/#", ... provides=[MessageSpec(name="value", dtype=DType.FLOAT)], ... params={"broker_host": "localhost"}, ... ) >>> await subscriber.start_subscription() """ if self.mode != MqttNode.Mode.SUBSCRIBER: raise RuntimeError("Can only start subscription for subscriber mode nodes") if self._subscription_task is not None: return # Already running self._stop_event = asyncio.Event() self._subscription_task = self._run_supervised("mqtt", self._subscription_loop) logger.info( "Started MQTT subscription for node '%s' to topics %s", self.name, list(self._topic_to_ports.keys()), ) async def stop_subscription(self) -> None: """ Stop the MQTT subscription. Gracefully stops the background subscription task. A supervised subscription is cancelled with the rest of them at teardown; only an unsupervised one is this method's to cancel. """ if self._stop_event is None: return self._stop_event.set() if self._subscription_task is not None: await self._cancel_task(self._subscription_task) self._subscription_task = None self._stop_event = None logger.info( "Stopped MQTT subscription for node '%s'", self.name, ) def _ports_for(self, topic: str) -> list[str]: """Which ports an arriving topic feeds — exact mapping, then filters.""" ports = self._topic_to_ports.get(topic) if ports is not None: return ports matched: list[str] = [] for pattern in self._wildcards: if topic_matches(pattern, topic): matched.extend(self._topic_to_ports[pattern]) return matched async def _subscription_loop(self) -> None: """ Listen for MQTT messages and trigger the pipeline. Subscribes to all unique topics from the ``topics`` mapping and uses the reverse lookup ``_topic_to_ports`` to route incoming payloads to the correct pipeline message names. One connection attempt: a dropped broker raises, and the supervisor decides when to try again. Reconnecting here as well would mean two backoff policies fighting over the same socket. """ import json import aiomqtt if not (self._stop_event and self._stop_event.is_set()): try: async with aiomqtt.Client( hostname=self.broker_host, port=self.broker_port, username=self.username, password=self.password, identifier=self.client_id, keepalive=self.keepalive, timeout=self.timeout, ) as client: # Subscribe to every unique topic for topic in self._topic_to_ports: await client.subscribe(topic, qos=self.qos) logger.info("[%s] Subscribed to %s", self.name, topic) self.report_health("ok") async for message in client.messages: if self._stop_event and self._stop_event.is_set(): break try: payload = message.payload.decode("utf-8") incoming_topic = str(message.topic) logger.info( "[%s] Received on %s: %s", self.name, incoming_topic, payload, ) # Find which port(s) this topic feeds ports = self._ports_for(incoming_topic) if not ports: logger.debug( "[%s] No mapping for topic '%s', ignoring", self.name, incoming_topic, ) continue # Parse the payload value try: parsed = json.loads(payload) except json.JSONDecodeError: parsed = payload by_port = {s.port: s for s in self.output_ports} typed_data = {} for port in ports: spec = by_port.get(port) if spec is None: continue # A JSON object may carry the value under a # named key — the port's own name, or whatever # ``json_key`` says the device wraps it in. key = self.json_keys.get(port, port) if isinstance(parsed, dict) and key in parsed: value = parsed[key] else: value = parsed typed_data[port] = spec.coerce(value) if typed_data: await asyncio.to_thread(self.inject, typed_data) except Exception as e: logger.error( "[%s] Error processing message: %s", self.name, e, exc_info=True, ) # Falling out of the message iterator without being told to # stop means the broker went away quietly. Raising is how the # supervisor hears about it. if not (self._stop_event and self._stop_event.is_set()): self.report_health("down", "subscription ended") raise ConnectionError( f"MQTT subscription for '{self.name}' ended unexpectedly" ) except asyncio.CancelledError: raise except Exception as e: logger.warning( "MQTT subscription for node '%s' failed: %s", self.name, e ) self.report_health("down", str(e)) if self._stop_event and self._stop_event.is_set(): return raise @property def is_subscribed(self) -> bool: """ Check if the subscription is currently active. :returns: True if subscription task is running. :rtype: bool """ return ( self._subscription_task is not None and not self._subscription_task.done() )