Split the node types into a package, and stop reconnecting per message
nodes.py had grown to 2k lines holding every integration behind a single blanket mypy exemption. It is now a package split by the outside world each node talks to, so the exemption shrinks to the four integration modules; base and mlp are type-checked, which turned up a dozen missing annotations. The senders opened a fresh connection — and, in the MQTT case, a fresh thread pool and event loop — for every single message. HTTP senders now share one pooled client, and a publisher holds one broker connection for its lifetime, fed from a bounded queue that drops the oldest value when the broker cannot keep up. An HTTP sender also no longer trips over a JSON reply that is not an object: outputs are keyed by port, so a bare scalar is a valid reply with nothing to publish. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_011LF61rxW1FG5YCD2J9YqjY
This commit is contained in:
+7
-3
@@ -60,9 +60,13 @@ Deferring because out of scope is fine, but don't mention deferring than.
|
|||||||
runs of that one job collide.
|
runs of that one job collide.
|
||||||
- PERF/UI: (deferred for now) the Monaco chunk is 2.6 MB. It only loads when a node panel opens, but the
|
- PERF/UI: (deferred for now) the Monaco chunk is 2.6 MB. It only loads when a node panel opens, but the
|
||||||
editor could be trimmed further or swapped for CodeMirror if that becomes a problem.
|
editor could be trimmed further or swapped for CodeMirror if that becomes a problem.
|
||||||
- CHORE/FLOW: `app/flow/nodes.py` is excluded from strict mypy (`[[tool.mypy.overrides]]` in
|
- CHORE/FLOW: four modules of `app/flow/nodes/` are still excluded from strict mypy
|
||||||
`pyproject.toml`). The node classes still carry prototype typing, `croniter` ships no stubs
|
(`[[tool.mypy.overrides]]` in `pyproject.toml`): `mqtt`, `http`, `influx` and `delay`. They
|
||||||
and `influxdb_client` does not re-export its names. Shrink it as each integration is revisited.
|
carry prototype typing, `croniter` ships no stubs and `influxdb_client` does not re-export
|
||||||
|
its names. `base` and `mlp` are checked; shrink the rest as each integration is revisited.
|
||||||
|
- CHORE/FLOW: a node function returning something other than a dict raises `AttributeError`
|
||||||
|
in `Node._to_messages` rather than a named error. Outputs are keyed by port, so a non-dict
|
||||||
|
cannot be one — say so where the return value is mapped.
|
||||||
- PERF/FLOW: every save rebuilds the whole pipeline. Fine at the current flow count; rebuild
|
- PERF/FLOW: every save rebuilds the whole pipeline. Fine at the current flow count; rebuild
|
||||||
only the touched flow when it starts to show.
|
only the touched flow when it starts to show.
|
||||||
- FEAT/UI: reintroduce `--chart-*` tokens as one designed sequential scale when the first
|
- FEAT/UI: reintroduce `--chart-*` tokens as one designed sequential scale when the first
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,21 @@
|
|||||||
|
"""Built-in node types.
|
||||||
|
|
||||||
|
Split by the outside world each one talks to. Importing from
|
||||||
|
``app.flow.nodes`` keeps working, which is what every caller does.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from app.flow.nodes.base import Node
|
||||||
|
from app.flow.nodes.delay import DelayNode
|
||||||
|
from app.flow.nodes.http import HttpNode
|
||||||
|
from app.flow.nodes.influx import InfluxDbNode
|
||||||
|
from app.flow.nodes.mlp import MLPNode
|
||||||
|
from app.flow.nodes.mqtt import MqttNode
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"DelayNode",
|
||||||
|
"HttpNode",
|
||||||
|
"InfluxDbNode",
|
||||||
|
"MLPNode",
|
||||||
|
"MqttNode",
|
||||||
|
"Node",
|
||||||
|
]
|
||||||
@@ -0,0 +1,305 @@
|
|||||||
|
"""The node base class: identity, ports, lifecycle and execution.
|
||||||
|
|
||||||
|
Every node type in this package derives from :class:`Node`; the type-specific
|
||||||
|
modules beside this one add what talking to a particular outside world means.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from collections.abc import Callable, Coroutine, Iterable
|
||||||
|
from typing import TYPE_CHECKING, Any, TypeAlias
|
||||||
|
|
||||||
|
from app.flow import logs
|
||||||
|
from app.flow.messages import MessageSpec, qualify
|
||||||
|
from app.flow.supervision import Supervisor
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from fastapi import FastAPI
|
||||||
|
|
||||||
|
from app.flow.pipeline import Pipeline
|
||||||
|
from app.flow.state import StateBackend
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# What a node hands back: the pipeline's state once it is bound, since a
|
||||||
|
# trigger runs the graph, and its own outputs when it is not.
|
||||||
|
NodeResult: TypeAlias = "StateBackend | dict[str, Any] | None"
|
||||||
|
|
||||||
|
|
||||||
|
class Node:
|
||||||
|
"""
|
||||||
|
A pipeline node that wraps a function with typed inputs/outputs.
|
||||||
|
|
||||||
|
Nodes are the fundamental building blocks of a pipeline. Each node encapsulates
|
||||||
|
a function that processes data, with explicit message-based inputs (requires)
|
||||||
|
and outputs (provides). Nodes can be connected into a directed acyclic graph (DAG)
|
||||||
|
where data flows from upstream to downstream nodes.
|
||||||
|
|
||||||
|
:param f: The function to execute when the node runs.
|
||||||
|
:type f: Callable
|
||||||
|
:param requires: Input messages this node consumes. Can be a single Message
|
||||||
|
or list of Messages. Empty for source nodes.
|
||||||
|
:type requires: MessageSpec | list[MessageSpec]
|
||||||
|
:param provides: Output messages this node produces. Can be a single Message
|
||||||
|
or list of Messages.
|
||||||
|
:type provides: MessageSpec | list[MessageSpec]
|
||||||
|
:param params: Additional parameters passed to the function during execution.
|
||||||
|
:type params: dict
|
||||||
|
:param name: Optional name for the node. Defaults to function name.
|
||||||
|
:type name: str | None
|
||||||
|
|
||||||
|
:ivar synchronous: If True, node only executes when all required inputs have
|
||||||
|
new versions since last execution. Useful for synchronizing multiple streams.
|
||||||
|
:vartype synchronous: bool
|
||||||
|
|
||||||
|
:example:
|
||||||
|
>>> def process_temp(temperature, params):
|
||||||
|
... return {"celsius": temperature * 0.5 + 32}
|
||||||
|
>>>
|
||||||
|
>>> temp_node = Node(
|
||||||
|
... f=process_temp,
|
||||||
|
... requires=MessageSpec(name="temperature", dtype=DType.FLOAT),
|
||||||
|
... provides=MessageSpec(name="celsius", dtype=DType.FLOAT),
|
||||||
|
... params={},
|
||||||
|
... )
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = (
|
||||||
|
"f",
|
||||||
|
"id",
|
||||||
|
"flow",
|
||||||
|
"name",
|
||||||
|
"input_ports",
|
||||||
|
"output_ports",
|
||||||
|
"requires",
|
||||||
|
"provides",
|
||||||
|
"params",
|
||||||
|
"_pipeline",
|
||||||
|
"synchronous",
|
||||||
|
"_on_health",
|
||||||
|
"supervisor",
|
||||||
|
)
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
f: Callable[..., Any],
|
||||||
|
requires: MessageSpec | Iterable[MessageSpec] = (),
|
||||||
|
provides: MessageSpec | Iterable[MessageSpec] = (),
|
||||||
|
params: dict[str, Any] | None = None,
|
||||||
|
name: str | None = None,
|
||||||
|
):
|
||||||
|
self.f = f
|
||||||
|
self._pipeline: Pipeline | None = None
|
||||||
|
self._on_health: Callable[[Node, str, str | None], None] | None = None
|
||||||
|
# Set by the controller before start(); a node built on its own (tests,
|
||||||
|
# previews) runs its loops unsupervised.
|
||||||
|
self.supervisor: Supervisor | None = None
|
||||||
|
self.params = dict(params) if params else {}
|
||||||
|
self.synchronous = bool(self.params.get("synchronous", False))
|
||||||
|
|
||||||
|
self.input_ports = self._normalize_ports(requires)
|
||||||
|
self.output_ports = self._normalize_ports(provides)
|
||||||
|
|
||||||
|
self.flow = ""
|
||||||
|
self.name: str = name or str(getattr(f, "__name__", "node"))
|
||||||
|
self.id: str = self.name
|
||||||
|
self._index_ports()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _normalize_ports(
|
||||||
|
msgs: MessageSpec | Iterable[MessageSpec],
|
||||||
|
) -> list[MessageSpec]:
|
||||||
|
"""Accept a single spec or any iterable of them."""
|
||||||
|
if isinstance(msgs, MessageSpec):
|
||||||
|
return [msgs]
|
||||||
|
return list(msgs or ())
|
||||||
|
|
||||||
|
def _index_ports(self) -> None:
|
||||||
|
"""Index bound ports by message name; unbound ports have no wiring."""
|
||||||
|
self.requires: dict[str, MessageSpec] = {
|
||||||
|
s.name: s for s in self.input_ports if s.name
|
||||||
|
}
|
||||||
|
self.provides: dict[str, MessageSpec] = {
|
||||||
|
s.name: s for s in self.output_ports if s.name
|
||||||
|
}
|
||||||
|
|
||||||
|
def assign_flow(self, flow: str, node_id: str) -> None:
|
||||||
|
"""Place this node in a flow, qualifying its identity and messages.
|
||||||
|
|
||||||
|
Called once by the loader, before the pipeline is built.
|
||||||
|
"""
|
||||||
|
self.flow = flow
|
||||||
|
self.name = node_id
|
||||||
|
self.id = f"{flow}.{node_id}"
|
||||||
|
self.input_ports = [
|
||||||
|
s.model_copy(update={"name": qualify(flow, s.name)})
|
||||||
|
for s in self.input_ports
|
||||||
|
]
|
||||||
|
self.output_ports = [
|
||||||
|
s.model_copy(update={"name": qualify(flow, s.name)})
|
||||||
|
for s in self.output_ports
|
||||||
|
]
|
||||||
|
self._index_ports()
|
||||||
|
|
||||||
|
@property
|
||||||
|
def local_id(self) -> str:
|
||||||
|
"""The node's id within its flow."""
|
||||||
|
return self.name
|
||||||
|
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
# Lifecycle
|
||||||
|
#
|
||||||
|
# A plain logic node has nothing to start or stop. The ones that talk to the
|
||||||
|
# outside — subscriptions, schedules, webhooks — override these, which is
|
||||||
|
# how the controller can bring a flow up or down without knowing what any
|
||||||
|
# particular node type is.
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
|
||||||
|
async def start(self, app: FastAPI | None = None) -> None:
|
||||||
|
"""Begin whatever this node listens to. Called when its flow starts."""
|
||||||
|
|
||||||
|
async def stop(self, app: FastAPI | None = None) -> None:
|
||||||
|
"""Undo :meth:`start`. Called before a rebuild, and must be idempotent."""
|
||||||
|
|
||||||
|
def _run_supervised(
|
||||||
|
self, name: str, factory: Callable[[], Coroutine[Any, Any, None]]
|
||||||
|
) -> asyncio.Task[None] | None:
|
||||||
|
"""Run a background loop under supervision where there is one.
|
||||||
|
|
||||||
|
Returns the task only when unsupervised, since that is the one case
|
||||||
|
where the caller has to cancel it itself.
|
||||||
|
"""
|
||||||
|
if self.supervisor is not None:
|
||||||
|
self.supervisor.spawn(f"{self.id}:{name}", self.flow, factory)
|
||||||
|
return None
|
||||||
|
return asyncio.create_task(factory())
|
||||||
|
|
||||||
|
def report_health(self, status: str, detail: str | None = None) -> None:
|
||||||
|
"""Say how this node's connection is doing: ok, degraded or down."""
|
||||||
|
if self._on_health is not None:
|
||||||
|
self._on_health(self, status, detail)
|
||||||
|
|
||||||
|
def bind(self, pipeline: Pipeline) -> None:
|
||||||
|
"""
|
||||||
|
Bind this node to a pipeline for external triggering.
|
||||||
|
|
||||||
|
Once bound, the node can trigger downstream execution when called.
|
||||||
|
This is typically done automatically during pipeline construction.
|
||||||
|
|
||||||
|
:param pipeline: The pipeline to bind this node to.
|
||||||
|
:type pipeline: Pipeline
|
||||||
|
"""
|
||||||
|
self._pipeline = pipeline
|
||||||
|
|
||||||
|
def __repr__(self) -> str:
|
||||||
|
return self.id
|
||||||
|
|
||||||
|
def __hash__(self) -> int:
|
||||||
|
return hash(self.id)
|
||||||
|
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
# Port translation
|
||||||
|
#
|
||||||
|
# The graph speaks qualified message names; node functions speak ports.
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
|
||||||
|
def _to_kwargs(self, inputs: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
kwargs = {}
|
||||||
|
for msg_name, value in inputs.items():
|
||||||
|
spec = self.requires.get(msg_name)
|
||||||
|
if spec is None:
|
||||||
|
continue
|
||||||
|
spec.check(value)
|
||||||
|
kwargs[spec.port] = value
|
||||||
|
return kwargs
|
||||||
|
|
||||||
|
def _to_messages(self, retval: dict[str, Any] | None) -> dict[str, Any] | None:
|
||||||
|
"""Map a function's port-keyed return value onto message names."""
|
||||||
|
if not retval:
|
||||||
|
return None
|
||||||
|
by_port = {s.port: s for s in self.output_ports if s.name}
|
||||||
|
outputs = {}
|
||||||
|
for key, value in retval.items():
|
||||||
|
spec = by_port.get(key) or self.provides.get(key)
|
||||||
|
if spec is None:
|
||||||
|
continue
|
||||||
|
spec.check(value)
|
||||||
|
outputs[spec.name] = value
|
||||||
|
return outputs or None
|
||||||
|
|
||||||
|
def execute(self, inputs: dict[str, Any] | None = None) -> dict[str, Any] | None:
|
||||||
|
"""Run the node function and return its outputs by message name.
|
||||||
|
|
||||||
|
Downstream nodes are not triggered — the pipeline schedules those.
|
||||||
|
"""
|
||||||
|
kwargs = self._to_kwargs(inputs or {})
|
||||||
|
return self._to_messages(self.f(**kwargs, params=self.params))
|
||||||
|
|
||||||
|
def trigger(self, inputs: dict[str, Any] | None = None) -> NodeResult:
|
||||||
|
"""
|
||||||
|
Trigger this node externally, executing downstream nodes if dependencies are met.
|
||||||
|
|
||||||
|
This method is for nodes that receive data from upstream dependencies.
|
||||||
|
For trigger/subscriber nodes that inject data into the pipeline, use :meth:`inject`.
|
||||||
|
|
||||||
|
:param inputs: Input values matching this node's ``requires``.
|
||||||
|
:type inputs: dict | None
|
||||||
|
:returns: Result of the node execution and downstream propagation.
|
||||||
|
:rtype: dict | None
|
||||||
|
:raises RuntimeError: If node is not bound to a pipeline.
|
||||||
|
"""
|
||||||
|
if self._pipeline is None:
|
||||||
|
raise RuntimeError("Node must be bound to a pipeline to trigger")
|
||||||
|
return self(inputs)
|
||||||
|
|
||||||
|
def inject(self, outputs: dict[str, Any] | None = None) -> NodeResult:
|
||||||
|
"""
|
||||||
|
Inject data into the pipeline as if this node produced it.
|
||||||
|
|
||||||
|
This method is for trigger/subscriber nodes that receive external data
|
||||||
|
(e.g., HTTP requests, MQTT messages) and need to inject it into the pipeline.
|
||||||
|
The data is validated against this node's ``provides`` specification.
|
||||||
|
|
||||||
|
For source nodes (nodes with no ``requires``), if no outputs are provided,
|
||||||
|
the node's function will be executed to generate outputs.
|
||||||
|
|
||||||
|
:param outputs: Output values matching this node's ``provides``.
|
||||||
|
:type outputs: dict | None
|
||||||
|
:returns: Result of downstream propagation.
|
||||||
|
:rtype: dict | None
|
||||||
|
:raises RuntimeError: If node is not bound to a pipeline.
|
||||||
|
:raises TypeError: If output values don't match ``provides`` types.
|
||||||
|
:raises KeyError: If required output keys are missing.
|
||||||
|
"""
|
||||||
|
if self._pipeline is None:
|
||||||
|
raise RuntimeError("Node must be bound to a pipeline to inject")
|
||||||
|
|
||||||
|
outputs = outputs or {}
|
||||||
|
|
||||||
|
# A source node asked to inject nothing produces its own data.
|
||||||
|
if not outputs and not self.requires:
|
||||||
|
collected = logs.Collector()
|
||||||
|
with logs.capture(collected):
|
||||||
|
outputs = self.f(params=self.params) or {}
|
||||||
|
self._pipeline.publish_log(self, collected, "")
|
||||||
|
|
||||||
|
return self._pipeline.trigger(self, self._to_messages(outputs))
|
||||||
|
|
||||||
|
def __call__(self, inputs: dict[str, Any] | None = None) -> NodeResult:
|
||||||
|
"""
|
||||||
|
Execute the node and trigger downstream nodes if bound to a pipeline.
|
||||||
|
|
||||||
|
Validates inputs against the node's ``requires`` specification, executes
|
||||||
|
the wrapped function, validates outputs, and triggers downstream execution
|
||||||
|
if the node is bound to a pipeline.
|
||||||
|
|
||||||
|
:param inputs: Input values keyed by message name. Must match the node's
|
||||||
|
``requires`` specification.
|
||||||
|
:type inputs: dict | None
|
||||||
|
:returns: Node outputs if successful, or pipeline execution results if bound.
|
||||||
|
:rtype: dict | None
|
||||||
|
"""
|
||||||
|
outputs = self.execute(inputs)
|
||||||
|
return self._pipeline.trigger(self, outputs) if self._pipeline else outputs
|
||||||
@@ -0,0 +1,323 @@
|
|||||||
|
"""Timing nodes: fixed delay, rate limit, and the cron schedule."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import time
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict
|
||||||
|
|
||||||
|
from app.flow.nodes.base import Node
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from fastapi import FastAPI
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class DelayNode(Node):
|
||||||
|
"""
|
||||||
|
A node that adds delay, rate-limiting, and/or cron-scheduled emissions.
|
||||||
|
|
||||||
|
This node can:
|
||||||
|
- Add a fixed delay before forwarding messages
|
||||||
|
- Rate-limit messages to a minimum interval between forwards
|
||||||
|
- Schedule emissions using crontab syntax (like a cron job)
|
||||||
|
|
||||||
|
The node passes through all input values to outputs with matching names.
|
||||||
|
If ``requires`` and ``provides`` have the same message names, the values are
|
||||||
|
forwarded directly. Otherwise, you can specify a ``mapping`` in params.
|
||||||
|
|
||||||
|
**Cron functionality:**
|
||||||
|
- Without ``requires``: Emits current timestamp at each cron tick
|
||||||
|
- With ``requires``: Emits last received input at each cron tick
|
||||||
|
|
||||||
|
**Execution order:** rate check → delay → forward
|
||||||
|
|
||||||
|
:param params: Parameters dict containing:
|
||||||
|
- ``delay`` (int): Fixed delay in seconds before forwarding (default: 0)
|
||||||
|
- ``interval`` (int): Minimum interval in seconds between forwards (default: 0)
|
||||||
|
- ``mapping`` (dict): Optional mapping from input names to output names
|
||||||
|
- ``cron`` (str): Crontab expression for scheduled emissions (optional).
|
||||||
|
Standard 5-field format: ``minute hour day-of-month month day-of-week``
|
||||||
|
:type params: dict
|
||||||
|
:param kwargs: Additional arguments passed to Node (requires, provides, name)
|
||||||
|
|
||||||
|
:example:
|
||||||
|
Simple passthrough with delay:
|
||||||
|
|
||||||
|
>>> delay_node = DelayNode(
|
||||||
|
... requires=[MessageSpec(name="temperature", dtype=DType.FLOAT)],
|
||||||
|
... provides=[MessageSpec(name="temperature", dtype=DType.FLOAT)],
|
||||||
|
... params={"delay": 1},
|
||||||
|
... )
|
||||||
|
|
||||||
|
Rate-limited forwarding:
|
||||||
|
|
||||||
|
>>> throttle_node = DelayNode(
|
||||||
|
... requires=[MessageSpec(name="sensor_data", dtype=DType.FLOAT)],
|
||||||
|
... provides=[MessageSpec(name="sensor_data", dtype=DType.FLOAT)],
|
||||||
|
... params={"interval": 5},
|
||||||
|
... )
|
||||||
|
|
||||||
|
Cron: emit timestamp every day at 08:30:
|
||||||
|
|
||||||
|
>>> alarm_node = DelayNode(
|
||||||
|
... provides=[MessageSpec(name="timestamp", dtype=DType.FLOAT)],
|
||||||
|
... params={"cron": "30 8 * * *"},
|
||||||
|
... )
|
||||||
|
|
||||||
|
Cron: emit stored input every 5 minutes:
|
||||||
|
|
||||||
|
>>> periodic_node = DelayNode(
|
||||||
|
... requires=[MessageSpec(name="value", dtype=DType.FLOAT)],
|
||||||
|
... provides=[MessageSpec(name="value", dtype=DType.FLOAT)],
|
||||||
|
... params={"cron": "*/5 * * * *"},
|
||||||
|
... )
|
||||||
|
|
||||||
|
Cron: emit on weekdays at 23:59:
|
||||||
|
|
||||||
|
>>> weekday_node = DelayNode(
|
||||||
|
... requires=[MessageSpec(name="daily_summary", dtype=DType.FLOAT)],
|
||||||
|
... provides=[MessageSpec(name="daily_summary", dtype=DType.FLOAT)],
|
||||||
|
... params={"cron": "59 23 * * 1-5"},
|
||||||
|
... )
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = (
|
||||||
|
"delay",
|
||||||
|
"interval",
|
||||||
|
"ts",
|
||||||
|
"mapping",
|
||||||
|
"cron_expr",
|
||||||
|
"last_input",
|
||||||
|
"_cron_task",
|
||||||
|
"_stop_cron",
|
||||||
|
)
|
||||||
|
|
||||||
|
class Params(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="allow")
|
||||||
|
|
||||||
|
delay: int = 0
|
||||||
|
interval: int = 0
|
||||||
|
# Input port → output port; ports are paired in order when empty.
|
||||||
|
mapping: dict[str, str] = {}
|
||||||
|
cron: str | None = None
|
||||||
|
|
||||||
|
def __init__(self, params: dict[str, Any] | None = None, **kwargs):
|
||||||
|
cfg = self.Params.model_validate(params or {})
|
||||||
|
self.delay = cfg.delay
|
||||||
|
self.interval = cfg.interval
|
||||||
|
self.mapping = cfg.mapping
|
||||||
|
self.cron_expr = cfg.cron
|
||||||
|
self.ts = 0.0
|
||||||
|
self.last_input: dict[str, Any] = {}
|
||||||
|
self._cron_task = None
|
||||||
|
self._stop_cron = None
|
||||||
|
|
||||||
|
super().__init__(f=self._f, params=params, **kwargs)
|
||||||
|
|
||||||
|
def _port_pairs(self) -> list[tuple[str, str]]:
|
||||||
|
"""Input port → output port pairs this node forwards along."""
|
||||||
|
if self.mapping:
|
||||||
|
return list(self.mapping.items())
|
||||||
|
return [
|
||||||
|
(i.port, o.port)
|
||||||
|
for i, o in zip(self.input_ports, self.output_ports, strict=False)
|
||||||
|
]
|
||||||
|
|
||||||
|
def _f(self, params, **kwargs):
|
||||||
|
"""Forward messages with optional delay, rate-limiting, and alarm."""
|
||||||
|
|
||||||
|
logger.info("[%s] Received %s", self.name, kwargs)
|
||||||
|
|
||||||
|
# Store last input for cron use
|
||||||
|
if kwargs:
|
||||||
|
self.last_input = dict(kwargs)
|
||||||
|
|
||||||
|
# Apply rate limiting first
|
||||||
|
if self.interval > 0:
|
||||||
|
ts = time.time()
|
||||||
|
if ts <= self.ts + self.interval:
|
||||||
|
logger.info("[%s] Stashing %s", self.name, kwargs)
|
||||||
|
return None
|
||||||
|
self.ts = ts
|
||||||
|
|
||||||
|
# Apply fixed delay
|
||||||
|
if self.delay > 0:
|
||||||
|
time.sleep(self.delay)
|
||||||
|
|
||||||
|
if not kwargs:
|
||||||
|
return None
|
||||||
|
|
||||||
|
output = {
|
||||||
|
out_port: kwargs[in_port]
|
||||||
|
for in_port, out_port in self._port_pairs()
|
||||||
|
if in_port in kwargs
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.info("[%s] Sending %s", self.name, output)
|
||||||
|
return output
|
||||||
|
|
||||||
|
# -----------------------------------------------------------------
|
||||||
|
# Cron scheduler
|
||||||
|
# -----------------------------------------------------------------
|
||||||
|
|
||||||
|
async def start(self, app: FastAPI | None = None) -> None:
|
||||||
|
"""Only a node with a schedule has anything to run on its own."""
|
||||||
|
if self.cron_expr:
|
||||||
|
await self.start_cron()
|
||||||
|
|
||||||
|
async def stop(self, app: FastAPI | None = None) -> None:
|
||||||
|
if self.cron_expr:
|
||||||
|
await self.stop_cron()
|
||||||
|
|
||||||
|
async def start_cron(self) -> None:
|
||||||
|
"""
|
||||||
|
Start the cron scheduler.
|
||||||
|
|
||||||
|
Requires the ``croniter`` package. The scheduler runs in the background
|
||||||
|
and triggers emissions according to the ``cron`` expression in params.
|
||||||
|
|
||||||
|
:raises ValueError: If no ``cron`` expression is configured.
|
||||||
|
"""
|
||||||
|
if not self.cron_expr:
|
||||||
|
logger.warning("No cron expression configured for node '%s'", self.name)
|
||||||
|
return
|
||||||
|
|
||||||
|
if self._cron_task is not None:
|
||||||
|
logger.info("Cron already running for node '%s'", self.name)
|
||||||
|
return
|
||||||
|
|
||||||
|
self._stop_cron = asyncio.Event()
|
||||||
|
self._cron_task = self._run_supervised("cron", self._cron_loop)
|
||||||
|
logger.info("Started cron for node '%s': %s", self.name, self.cron_expr)
|
||||||
|
|
||||||
|
async def stop_cron(self) -> None:
|
||||||
|
"""Stop the cron scheduler."""
|
||||||
|
if self._stop_cron is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
self._stop_cron.set()
|
||||||
|
|
||||||
|
if self._cron_task is not None:
|
||||||
|
self._cron_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._cron_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
self._cron_task = None
|
||||||
|
self._stop_cron = None
|
||||||
|
logger.info("Stopped cron for node '%s'", self.name)
|
||||||
|
|
||||||
|
async def _cron_loop(self):
|
||||||
|
"""Background loop that sleeps until the next cron tick and triggers."""
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
try:
|
||||||
|
from croniter import croniter
|
||||||
|
except ImportError:
|
||||||
|
logger.error(
|
||||||
|
"croniter package is required for cron scheduling. "
|
||||||
|
"Install it with: pip install croniter"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if not croniter.is_valid(self.cron_expr):
|
||||||
|
logger.error(
|
||||||
|
"Invalid cron expression '%s' for node '%s'",
|
||||||
|
self.cron_expr,
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
cron = croniter(self.cron_expr, datetime.now())
|
||||||
|
|
||||||
|
while not (self._stop_cron and self._stop_cron.is_set()):
|
||||||
|
try:
|
||||||
|
# Compute seconds until next tick
|
||||||
|
next_dt = cron.get_next(datetime)
|
||||||
|
now = datetime.now()
|
||||||
|
wait_seconds = max(0, (next_dt - now).total_seconds())
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Node '%s' cron: next tick at %s (in %.1fs)",
|
||||||
|
self.name,
|
||||||
|
next_dt.isoformat(),
|
||||||
|
wait_seconds,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Sleep until next tick (wake up on stop signal)
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(self._stop_cron.wait(), timeout=wait_seconds)
|
||||||
|
# If we get here, stop was requested
|
||||||
|
break
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
# Timeout means it's time to fire
|
||||||
|
pass
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Cron triggered for node '%s' (%s)",
|
||||||
|
self.name,
|
||||||
|
self.cron_expr,
|
||||||
|
)
|
||||||
|
await self._trigger_cron()
|
||||||
|
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception:
|
||||||
|
# One bad tick should not cost the schedule; the supervisor
|
||||||
|
# picks it up from here, backoff included.
|
||||||
|
logger.error(
|
||||||
|
"Error in cron loop for node '%s'", self.name, exc_info=True
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def _trigger_cron(self):
|
||||||
|
"""Emit data into the pipeline on a cron tick."""
|
||||||
|
if self._pipeline is None:
|
||||||
|
logger.warning(
|
||||||
|
"Node '%s' not bound to pipeline, cannot trigger cron",
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
output = {}
|
||||||
|
|
||||||
|
if not self.input_ports:
|
||||||
|
# Nothing to forward — emit the current time on every output.
|
||||||
|
output = {spec.port: time.time() for spec in self.output_ports}
|
||||||
|
logger.info("Cron emission for node '%s': timestamp", self.name)
|
||||||
|
elif self.last_input:
|
||||||
|
output = {
|
||||||
|
out_port: self.last_input[in_port]
|
||||||
|
for in_port, out_port in self._port_pairs()
|
||||||
|
if in_port in self.last_input
|
||||||
|
}
|
||||||
|
logger.info(
|
||||||
|
"Cron emission for node '%s': stored input %s",
|
||||||
|
self.name,
|
||||||
|
output,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Cron for node '%s' triggered but no input stored yet",
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if output:
|
||||||
|
await asyncio.to_thread(
|
||||||
|
self._pipeline.trigger, self, self._to_messages(output)
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
"Error triggering cron for node '%s': %s",
|
||||||
|
self.name,
|
||||||
|
e,
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
@@ -0,0 +1,401 @@
|
|||||||
|
"""HTTP nodes: an inbound webhook trigger and an outbound request sender."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hmac
|
||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
from collections.abc import Iterable
|
||||||
|
from enum import Enum
|
||||||
|
from typing import TYPE_CHECKING, Any, Literal
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
|
|
||||||
|
from app.flow.messages import MessageSpec
|
||||||
|
from app.flow.nodes.base import Node
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from fastapi import FastAPI
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# One pooled client for every sender node: connections are the expensive part
|
||||||
|
# of an HTTP request, and a node that fires every second should keep its own.
|
||||||
|
_client: httpx.Client | None = None
|
||||||
|
_client_lock = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
def shared_client() -> httpx.Client:
|
||||||
|
"""The process-wide HTTP client, built on first use."""
|
||||||
|
global _client
|
||||||
|
with _client_lock:
|
||||||
|
if _client is None:
|
||||||
|
_client = httpx.Client()
|
||||||
|
return _client
|
||||||
|
|
||||||
|
|
||||||
|
def close_shared_client() -> None:
|
||||||
|
"""Release pooled connections at shutdown. Safe to call more than once."""
|
||||||
|
global _client
|
||||||
|
with _client_lock:
|
||||||
|
if _client is not None:
|
||||||
|
_client.close()
|
||||||
|
_client = None
|
||||||
|
|
||||||
|
|
||||||
|
class HttpNode(Node):
|
||||||
|
"""
|
||||||
|
HTTP node that can act as a trigger (receiver) or sender based on configuration.
|
||||||
|
|
||||||
|
This node integrates with FastAPI to either:
|
||||||
|
|
||||||
|
- **Trigger mode**: Receive incoming HTTP requests (GET/POST) and inject data
|
||||||
|
into the pipeline. Used when ``provides`` is specified but ``requires`` is empty.
|
||||||
|
- **Sender mode**: Make outgoing HTTP requests with pipeline data. Used when
|
||||||
|
``requires`` is specified.
|
||||||
|
|
||||||
|
:param url: The URL endpoint. For trigger mode, this is the route path
|
||||||
|
(e.g., "/sensors/temperature"). For sender mode, this is the full URL
|
||||||
|
to send requests to.
|
||||||
|
:type url: str
|
||||||
|
:param method: HTTP method - "GET" or "POST".
|
||||||
|
:type method: Literal["GET", "POST"]
|
||||||
|
:param requires: Messages required by this node (makes it a sender node).
|
||||||
|
:type requires: MessageSpec | list[MessageSpec] | None
|
||||||
|
:param provides: Messages provided by this node (makes it a trigger node).
|
||||||
|
:type provides: MessageSpec | list[MessageSpec] | None
|
||||||
|
:param params: Additional parameters for the node.
|
||||||
|
:type params: dict
|
||||||
|
:param name: Optional name for the node.
|
||||||
|
:type name: str | None
|
||||||
|
:param timeout: Request timeout in seconds (for sender mode).
|
||||||
|
:type timeout: float
|
||||||
|
:param headers: Additional HTTP headers.
|
||||||
|
:type headers: dict[str, str] | None
|
||||||
|
:param secret: Shared secret callers append to the webhook URL (trigger mode).
|
||||||
|
:type secret: str
|
||||||
|
|
||||||
|
:raises ValueError: If both ``requires`` and ``provides`` are empty, or if
|
||||||
|
the configuration is invalid.
|
||||||
|
|
||||||
|
:example:
|
||||||
|
Trigger node (receives POST requests):
|
||||||
|
|
||||||
|
>>> trigger = HttpNode(
|
||||||
|
... url="/api/sensors/temperature",
|
||||||
|
... method="POST",
|
||||||
|
... provides=[MessageSpec(name="temperature", dtype=DType.FLOAT)],
|
||||||
|
... params={},
|
||||||
|
... )
|
||||||
|
|
||||||
|
Sender node (makes POST requests):
|
||||||
|
|
||||||
|
>>> sender = HttpNode(
|
||||||
|
... url="https://api.example.com/data",
|
||||||
|
... method="POST",
|
||||||
|
... requires=[MessageSpec(name="temperature", dtype=DType.FLOAT)],
|
||||||
|
... params={},
|
||||||
|
... )
|
||||||
|
"""
|
||||||
|
|
||||||
|
class Mode(Enum):
|
||||||
|
"""Operating mode of the HTTP node."""
|
||||||
|
|
||||||
|
TRIGGER = "trigger" # Receives HTTP requests
|
||||||
|
SENDER = "sender" # Sends HTTP requests
|
||||||
|
|
||||||
|
__slots__ = (
|
||||||
|
"url",
|
||||||
|
"method",
|
||||||
|
"mode",
|
||||||
|
"timeout",
|
||||||
|
"headers",
|
||||||
|
"secret",
|
||||||
|
"_route_registered",
|
||||||
|
)
|
||||||
|
|
||||||
|
class Params(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="allow")
|
||||||
|
|
||||||
|
url: str
|
||||||
|
method: Literal["GET", "POST"] = "POST"
|
||||||
|
timeout: float = 30.0
|
||||||
|
headers: dict[str, str] = {}
|
||||||
|
secret: str = Field(
|
||||||
|
default="",
|
||||||
|
description=(
|
||||||
|
"Shared secret for a webhook, appended to its URL: "
|
||||||
|
"/hooks/<flow>/<url>/<secret>. Empty leaves the webhook open "
|
||||||
|
"to anyone."
|
||||||
|
),
|
||||||
|
json_schema_extra={"x-secret": True},
|
||||||
|
)
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
url: str | None = None,
|
||||||
|
method: Literal["GET", "POST"] | None = None,
|
||||||
|
requires: MessageSpec | Iterable[MessageSpec] = (),
|
||||||
|
provides: MessageSpec | Iterable[MessageSpec] = (),
|
||||||
|
params: dict[str, Any] | None = None,
|
||||||
|
name: str | None = None,
|
||||||
|
):
|
||||||
|
params = dict(params) if params else {}
|
||||||
|
if url is not None:
|
||||||
|
params["url"] = url
|
||||||
|
if method is not None:
|
||||||
|
params["method"] = method
|
||||||
|
cfg = self.Params.model_validate(params)
|
||||||
|
|
||||||
|
requires = Node._normalize_ports(requires)
|
||||||
|
provides = Node._normalize_ports(provides)
|
||||||
|
if not requires and not provides:
|
||||||
|
raise ValueError(
|
||||||
|
"An HTTP node needs either inputs (to send) or outputs (to receive)"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Inputs mean this node sends data out; outputs mean it receives.
|
||||||
|
self.mode = HttpNode.Mode.SENDER if requires else HttpNode.Mode.TRIGGER
|
||||||
|
|
||||||
|
self.url = cfg.url
|
||||||
|
self.method = cfg.method.upper()
|
||||||
|
self.timeout = cfg.timeout
|
||||||
|
self.headers = cfg.headers
|
||||||
|
self.secret = cfg.secret
|
||||||
|
self._route_registered = False
|
||||||
|
|
||||||
|
# Set default name based on mode and URL
|
||||||
|
if name is None:
|
||||||
|
safe_url = self.url.replace("/", "_").replace(":", "").strip("_")
|
||||||
|
name = f"http_{self.mode.value}_{safe_url}"
|
||||||
|
|
||||||
|
# Initialize parent with appropriate function
|
||||||
|
# For trigger mode, f is a no-op since data is injected via inject()
|
||||||
|
# For sender mode, f handles the outgoing HTTP request
|
||||||
|
super().__init__(
|
||||||
|
f=(
|
||||||
|
self._noop_trigger
|
||||||
|
if self.mode == HttpNode.Mode.TRIGGER
|
||||||
|
else self._sender_handler
|
||||||
|
),
|
||||||
|
requires=requires,
|
||||||
|
provides=provides,
|
||||||
|
params=params,
|
||||||
|
name=name,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _noop_trigger(params: dict, **kwargs) -> dict | None:
|
||||||
|
"""
|
||||||
|
No-op function for trigger mode nodes.
|
||||||
|
|
||||||
|
Trigger mode nodes inject data via :meth:`inject`, not :meth:`__call__`.
|
||||||
|
This function exists only to satisfy the Node interface.
|
||||||
|
"""
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _sender_handler(self, params: dict, **kwargs) -> dict | None:
|
||||||
|
"""
|
||||||
|
Send HTTP request with pipeline data (sender mode).
|
||||||
|
|
||||||
|
This method is called when upstream dependencies are satisfied.
|
||||||
|
It sends the required data via HTTP request.
|
||||||
|
|
||||||
|
Node functions already run on a worker thread, so the request is sent
|
||||||
|
synchronously over the shared pooled client: a node publishing every
|
||||||
|
second should reuse its connection, not open one per message.
|
||||||
|
|
||||||
|
:param params: Node parameters.
|
||||||
|
:type params: dict
|
||||||
|
:param kwargs: Pipeline data to send (from required messages).
|
||||||
|
:type kwargs: Any
|
||||||
|
:returns: Response data if the endpoint returns JSON, None otherwise.
|
||||||
|
:rtype: dict | None
|
||||||
|
"""
|
||||||
|
client = shared_client()
|
||||||
|
try:
|
||||||
|
if self.method == "GET":
|
||||||
|
response = client.get(
|
||||||
|
self.url,
|
||||||
|
params=kwargs,
|
||||||
|
headers=self.headers,
|
||||||
|
timeout=self.timeout,
|
||||||
|
)
|
||||||
|
else: # POST
|
||||||
|
response = client.post(
|
||||||
|
self.url,
|
||||||
|
json=kwargs,
|
||||||
|
headers=self.headers,
|
||||||
|
timeout=self.timeout,
|
||||||
|
)
|
||||||
|
|
||||||
|
response.raise_for_status()
|
||||||
|
|
||||||
|
try:
|
||||||
|
body = response.json()
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
# Outputs are keyed by port, so only an object can be one. A bare
|
||||||
|
# scalar or list is a valid reply, just not something to publish.
|
||||||
|
return body if isinstance(body, dict) else None
|
||||||
|
|
||||||
|
except httpx.HTTPStatusError as e:
|
||||||
|
logger.error(
|
||||||
|
"HTTP error in node '%s': status=%s, url=%s",
|
||||||
|
self.name,
|
||||||
|
e.response.status_code,
|
||||||
|
self.url,
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
except httpx.RequestError as e:
|
||||||
|
logger.error("Request error in node '%s': %s", self.name, e)
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def start(self, app: FastAPI | None = None) -> None:
|
||||||
|
"""A webhook needs a route; a sender reaches out on its own."""
|
||||||
|
if self.mode is HttpNode.Mode.TRIGGER and app is not None:
|
||||||
|
self.register_route(app)
|
||||||
|
|
||||||
|
async def stop(self, app: FastAPI | None = None) -> None:
|
||||||
|
if self.mode is HttpNode.Mode.TRIGGER and app is not None:
|
||||||
|
self.unregister_route(app)
|
||||||
|
|
||||||
|
def register_route(self, app: FastAPI) -> None:
|
||||||
|
"""
|
||||||
|
Register this node's HTTP endpoint with a FastAPI application.
|
||||||
|
|
||||||
|
This method should only be called for trigger mode nodes. It creates
|
||||||
|
a route that, when called, triggers the node in the pipeline.
|
||||||
|
|
||||||
|
:param app: The FastAPI application instance.
|
||||||
|
:type app: FastAPI
|
||||||
|
:raises RuntimeError: If called on a sender mode node.
|
||||||
|
|
||||||
|
:example:
|
||||||
|
>>> from fastapi import FastAPI
|
||||||
|
>>> app = FastAPI()
|
||||||
|
>>> trigger_node = HttpNode(
|
||||||
|
... url="/sensors/data",
|
||||||
|
... method="POST",
|
||||||
|
... provides=[MessageSpec(name="value", dtype=DType.FLOAT)],
|
||||||
|
... params={},
|
||||||
|
... )
|
||||||
|
>>> trigger_node.register_route(app)
|
||||||
|
"""
|
||||||
|
if self.mode != HttpNode.Mode.TRIGGER:
|
||||||
|
raise RuntimeError("Can only register routes for trigger mode nodes")
|
||||||
|
|
||||||
|
if self._route_registered:
|
||||||
|
return
|
||||||
|
|
||||||
|
from starlette.requests import Request
|
||||||
|
from starlette.responses import JSONResponse
|
||||||
|
from starlette.routing import Route
|
||||||
|
|
||||||
|
async def handle_request(request: Request) -> JSONResponse:
|
||||||
|
"""
|
||||||
|
Handle incoming HTTP request and trigger the pipeline.
|
||||||
|
|
||||||
|
:param request: The incoming Starlette request.
|
||||||
|
:type request: Request
|
||||||
|
:returns: JSON response with trigger result.
|
||||||
|
:rtype: JSONResponse
|
||||||
|
"""
|
||||||
|
if self.secret and not hmac.compare_digest(
|
||||||
|
str(request.path_params.get("secret", "")).encode(),
|
||||||
|
self.secret.encode(),
|
||||||
|
):
|
||||||
|
# A wrong secret must look like no such hook at all.
|
||||||
|
return JSONResponse(content={"detail": "Not Found"}, status_code=404)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Parse request data
|
||||||
|
if self.method == "GET":
|
||||||
|
data = dict(request.query_params)
|
||||||
|
else: # POST
|
||||||
|
content_type = request.headers.get("content-type", "")
|
||||||
|
if "application/json" in content_type:
|
||||||
|
data = await request.json()
|
||||||
|
elif "application/x-www-form-urlencoded" in content_type:
|
||||||
|
form = await request.form()
|
||||||
|
data = dict(form)
|
||||||
|
else:
|
||||||
|
data = await request.json() # Default to JSON
|
||||||
|
|
||||||
|
# Payload keys are port names; values arrive as text.
|
||||||
|
typed_data = {}
|
||||||
|
for spec in self.output_ports:
|
||||||
|
if spec.port in data:
|
||||||
|
typed_data[spec.port] = spec.coerce(data[spec.port])
|
||||||
|
|
||||||
|
# Inject data into the pipeline (trigger mode nodes inject via provides)
|
||||||
|
result = self.inject(typed_data)
|
||||||
|
|
||||||
|
return JSONResponse(
|
||||||
|
content={
|
||||||
|
"status": "triggered",
|
||||||
|
"node": self.name,
|
||||||
|
"data": typed_data,
|
||||||
|
"result": result if isinstance(result, dict) else None,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
except ValueError as e:
|
||||||
|
return JSONResponse(
|
||||||
|
content={"error": str(e)},
|
||||||
|
status_code=400,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
return JSONResponse(
|
||||||
|
content={"error": str(e)},
|
||||||
|
status_code=500,
|
||||||
|
)
|
||||||
|
|
||||||
|
# The secret is a path parameter, not part of the registered path, so
|
||||||
|
# neither the route table nor the logs below carry its value.
|
||||||
|
path = f"{self.url}/{{secret}}" if self.secret else self.url
|
||||||
|
|
||||||
|
# Create a Starlette Route and add it directly to the app's routes
|
||||||
|
route = Route(
|
||||||
|
path,
|
||||||
|
handle_request,
|
||||||
|
methods=[self.method],
|
||||||
|
name=self.id,
|
||||||
|
)
|
||||||
|
app.routes.append(route)
|
||||||
|
|
||||||
|
self._route_registered = True
|
||||||
|
logger.info(
|
||||||
|
"Registered %s route for node '%s': %s",
|
||||||
|
self.method,
|
||||||
|
self.name,
|
||||||
|
path,
|
||||||
|
)
|
||||||
|
|
||||||
|
def unregister_route(self, app: FastAPI) -> None:
|
||||||
|
"""
|
||||||
|
Unregister this node's HTTP endpoint from a FastAPI application.
|
||||||
|
|
||||||
|
:param app: The FastAPI application instance.
|
||||||
|
:type app: FastAPI
|
||||||
|
|
||||||
|
.. note::
|
||||||
|
FastAPI doesn't natively support route removal. This method
|
||||||
|
removes the route from the internal routes list, but the change
|
||||||
|
may not take effect until the application is restarted or
|
||||||
|
the OpenAPI schema is regenerated.
|
||||||
|
"""
|
||||||
|
if not self._route_registered:
|
||||||
|
return
|
||||||
|
|
||||||
|
# FastAPI doesn't have a clean way to remove routes
|
||||||
|
# We need to filter them out from the routes list
|
||||||
|
app.routes[:] = [
|
||||||
|
route
|
||||||
|
for route in app.routes
|
||||||
|
if not (hasattr(route, "name") and route.name == self.id)
|
||||||
|
]
|
||||||
|
|
||||||
|
self._route_registered = False
|
||||||
|
logger.info("Unregistered route for node '%s': %s", self.name, self.url)
|
||||||
@@ -0,0 +1,495 @@
|
|||||||
|
"""InfluxDB nodes: write points from messages, or read a query into them."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from collections.abc import Iterable
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
|
|
||||||
|
from app.flow.messages import MessageSpec
|
||||||
|
from app.flow.nodes.base import Node
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class InfluxDbNode(Node):
|
||||||
|
"""
|
||||||
|
InfluxDB node for writing to and reading from InfluxDB.
|
||||||
|
|
||||||
|
This node can perform both write and read operations independently:
|
||||||
|
|
||||||
|
- **Write operation**: Triggered when upstream dependencies are satisfied
|
||||||
|
(data flows in via ``requires``). Writes data points to InfluxDB based
|
||||||
|
on the ``writes`` configuration in params.
|
||||||
|
- **Read operation**: Performed when the node provides data to downstream
|
||||||
|
nodes via ``provides``, based on the ``queries`` configuration in params.
|
||||||
|
|
||||||
|
Both operations use a similar configuration pattern in params, making the
|
||||||
|
API consistent and the input/output data simple (just values).
|
||||||
|
|
||||||
|
:param requires: Messages to write to InfluxDB. The actual value from each
|
||||||
|
message is written according to the corresponding config in ``writes``.
|
||||||
|
:type requires: MessageSpec | list[MessageSpec] | None
|
||||||
|
:param provides: Messages to read from InfluxDB. Each message gets its value
|
||||||
|
from a query defined in ``queries``.
|
||||||
|
:type provides: MessageSpec | list[MessageSpec] | None
|
||||||
|
:param params: Parameters dict containing:
|
||||||
|
- ``url`` (str): InfluxDB server URL (required)
|
||||||
|
- ``token`` (str): Authentication token (required)
|
||||||
|
- ``org`` (str): Organization name (required)
|
||||||
|
- ``bucket`` (str): Bucket name (required)
|
||||||
|
- ``write_precision`` (str): Write precision ("ns", "us", "ms", "s"), default "ms"
|
||||||
|
- ``query_range`` (str): Default time range for queries, e.g., "-1h", "-24h"
|
||||||
|
- ``writes`` (dict): Write configurations keyed by message name, each with:
|
||||||
|
- ``measurement`` (str): Measurement name to write to
|
||||||
|
- ``field`` (str): Field name to write (default: "value")
|
||||||
|
- ``tags`` (dict): Static tags to add to each point
|
||||||
|
- ``queries`` (dict): Query configurations keyed by message name, each with:
|
||||||
|
- ``measurement`` (str): Measurement name to query
|
||||||
|
- ``field`` (str): Field name to retrieve (default: "value")
|
||||||
|
- ``tags`` (dict): Optional tag filters
|
||||||
|
- ``range`` (str): Optional time range override
|
||||||
|
- ``aggregation`` (str): Aggregation function ("mean", "last", "first", "max", "min")
|
||||||
|
:type params: dict
|
||||||
|
:param name: Optional name for the node.
|
||||||
|
:type name: str | None
|
||||||
|
|
||||||
|
:raises ValueError: If required params are missing or both requires and provides are empty.
|
||||||
|
|
||||||
|
:example:
|
||||||
|
Write-only node (writes temperature values):
|
||||||
|
|
||||||
|
>>> writer = InfluxDbNode(
|
||||||
|
... requires=[MessageSpec(name="temperature", dtype=DType.FLOAT)],
|
||||||
|
... params={
|
||||||
|
... "url": "http://localhost:8086",
|
||||||
|
... "token": "my-token",
|
||||||
|
... "org": "my-org",
|
||||||
|
... "bucket": "sensors",
|
||||||
|
... "writes": {
|
||||||
|
... "temperature": {
|
||||||
|
... "measurement": "environment",
|
||||||
|
... "field": "temp_celsius",
|
||||||
|
... "tags": {"location": "room1", "sensor": "dht22"},
|
||||||
|
... }
|
||||||
|
... },
|
||||||
|
... },
|
||||||
|
... )
|
||||||
|
|
||||||
|
Read-only node (queries average temperature):
|
||||||
|
|
||||||
|
>>> reader = InfluxDbNode(
|
||||||
|
... provides=[MessageSpec(name="avg_temperature", dtype=DType.FLOAT)],
|
||||||
|
... params={
|
||||||
|
... "url": "http://localhost:8086",
|
||||||
|
... "token": "my-token",
|
||||||
|
... "org": "my-org",
|
||||||
|
... "bucket": "sensors",
|
||||||
|
... "queries": {
|
||||||
|
... "avg_temperature": {
|
||||||
|
... "measurement": "environment",
|
||||||
|
... "field": "temp_celsius",
|
||||||
|
... "tags": {"location": "room1"},
|
||||||
|
... "range": "-1h",
|
||||||
|
... "aggregation": "mean",
|
||||||
|
... }
|
||||||
|
... },
|
||||||
|
... },
|
||||||
|
... )
|
||||||
|
|
||||||
|
Combined read/write node:
|
||||||
|
|
||||||
|
>>> node = InfluxDbNode(
|
||||||
|
... requires=[MessageSpec(name="raw_temp", dtype=DType.FLOAT)],
|
||||||
|
... provides=[MessageSpec(name="avg_temp", dtype=DType.FLOAT)],
|
||||||
|
... params={
|
||||||
|
... "url": "http://localhost:8086",
|
||||||
|
... "token": "my-token",
|
||||||
|
... "org": "my-org",
|
||||||
|
... "bucket": "sensors",
|
||||||
|
... "writes": {
|
||||||
|
... "raw_temp": {
|
||||||
|
... "measurement": "temperature",
|
||||||
|
... "field": "value",
|
||||||
|
... "tags": {"source": "sensor"},
|
||||||
|
... }
|
||||||
|
... },
|
||||||
|
... "queries": {
|
||||||
|
... "avg_temp": {
|
||||||
|
... "measurement": "temperature",
|
||||||
|
... "field": "value",
|
||||||
|
... "aggregation": "mean",
|
||||||
|
... "range": "-5m",
|
||||||
|
... }
|
||||||
|
... },
|
||||||
|
... },
|
||||||
|
... )
|
||||||
|
"""
|
||||||
|
|
||||||
|
__slots__ = (
|
||||||
|
"url",
|
||||||
|
"token",
|
||||||
|
"org",
|
||||||
|
"bucket",
|
||||||
|
"write_precision",
|
||||||
|
"query_range",
|
||||||
|
"writes",
|
||||||
|
"queries",
|
||||||
|
"_write_client",
|
||||||
|
"_query_client",
|
||||||
|
)
|
||||||
|
|
||||||
|
class Params(BaseModel):
|
||||||
|
model_config = ConfigDict(extra="allow")
|
||||||
|
|
||||||
|
url: str
|
||||||
|
token: str = Field(json_schema_extra={"x-secret": True})
|
||||||
|
org: str
|
||||||
|
bucket: str
|
||||||
|
write_precision: str = "ms"
|
||||||
|
query_range: str = "-1h"
|
||||||
|
# Per-port write and query configuration.
|
||||||
|
writes: dict[str, dict[str, Any]] = {}
|
||||||
|
queries: dict[str, dict[str, Any]] = {}
|
||||||
|
|
||||||
|
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 InfluxDB node needs either inputs (to write) or outputs (to read)"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.url = cfg.url
|
||||||
|
self.token = cfg.token
|
||||||
|
self.org = cfg.org
|
||||||
|
self.bucket = cfg.bucket
|
||||||
|
self.write_precision = cfg.write_precision
|
||||||
|
self.query_range = cfg.query_range
|
||||||
|
self.writes = cfg.writes
|
||||||
|
self.queries = cfg.queries
|
||||||
|
|
||||||
|
# Lazy-initialized clients
|
||||||
|
self._write_client = None
|
||||||
|
self._query_client = None
|
||||||
|
|
||||||
|
# Set default name
|
||||||
|
if name is None:
|
||||||
|
name = f"influxdb_{self.bucket}"
|
||||||
|
|
||||||
|
# Initialize parent
|
||||||
|
# The handler function depends on what operations are configured
|
||||||
|
super().__init__(
|
||||||
|
f=self._handler,
|
||||||
|
requires=requires,
|
||||||
|
provides=provides,
|
||||||
|
params=params,
|
||||||
|
name=name,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _handler(self, params: dict, **kwargs) -> dict | None:
|
||||||
|
"""
|
||||||
|
Handle incoming data - write to InfluxDB and optionally query.
|
||||||
|
|
||||||
|
This method is called when upstream dependencies (requires) are satisfied.
|
||||||
|
It writes the incoming data to InfluxDB and can also perform reads.
|
||||||
|
|
||||||
|
:param params: Node parameters.
|
||||||
|
:type params: dict
|
||||||
|
:param kwargs: Incoming data from upstream nodes.
|
||||||
|
:type kwargs: Any
|
||||||
|
:returns: Query results if provides is configured, None otherwise.
|
||||||
|
:rtype: dict | None
|
||||||
|
"""
|
||||||
|
# Write incoming data
|
||||||
|
if kwargs:
|
||||||
|
self._write_points(kwargs)
|
||||||
|
|
||||||
|
# If we have provides, perform queries
|
||||||
|
if self.provides:
|
||||||
|
return self._query_data()
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _write_points(self, data: dict) -> None:
|
||||||
|
"""
|
||||||
|
Write data points to InfluxDB using configuration from ``writes``.
|
||||||
|
|
||||||
|
The write configuration is looked up in ``self.writes`` by message name.
|
||||||
|
Each config specifies the measurement, field, and tags. The actual value
|
||||||
|
comes from the incoming data.
|
||||||
|
|
||||||
|
:param data: Data to write, keyed by port name. Values can be:
|
||||||
|
- Simple values (float, int, str, bool): Written using config from ``writes``
|
||||||
|
- List of values: Each value written as a separate point
|
||||||
|
- Dict with "value" key: Value extracted and written using config
|
||||||
|
- Dict with "value" and "tags" keys: Value written with merged tags
|
||||||
|
:type data: dict
|
||||||
|
"""
|
||||||
|
from influxdb_client import InfluxDBClient, Point, WritePrecision
|
||||||
|
from influxdb_client.client.write_api import SYNCHRONOUS
|
||||||
|
|
||||||
|
try:
|
||||||
|
with InfluxDBClient(url=self.url, token=self.token, org=self.org) as client:
|
||||||
|
write_api = client.write_api(write_options=SYNCHRONOUS)
|
||||||
|
|
||||||
|
precision_map = {
|
||||||
|
"ns": WritePrecision.NS,
|
||||||
|
"us": WritePrecision.US,
|
||||||
|
"ms": WritePrecision.MS,
|
||||||
|
"s": WritePrecision.S,
|
||||||
|
}
|
||||||
|
precision = precision_map.get(self.write_precision, WritePrecision.MS)
|
||||||
|
|
||||||
|
for msg_name, msg_value in data.items():
|
||||||
|
# Get write configuration for this message
|
||||||
|
write_config = self.writes.get(msg_name, {})
|
||||||
|
|
||||||
|
# Get measurement, field, and base tags from config
|
||||||
|
measurement = write_config.get("measurement", msg_name)
|
||||||
|
field = write_config.get("field", "value")
|
||||||
|
base_tags = write_config.get("tags", {})
|
||||||
|
|
||||||
|
# Handle list of values (batch write)
|
||||||
|
values_to_write = (
|
||||||
|
msg_value if isinstance(msg_value, list) else [msg_value]
|
||||||
|
)
|
||||||
|
|
||||||
|
for item in values_to_write:
|
||||||
|
# Extract value and optional runtime tags
|
||||||
|
if isinstance(item, dict):
|
||||||
|
value = item.get("value", item)
|
||||||
|
runtime_tags = item.get("tags", {})
|
||||||
|
# If no "value" key, treat the whole dict as invalid
|
||||||
|
if "value" not in item and not isinstance(
|
||||||
|
value, (int, float, str, bool)
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
"Skipping invalid item in node '%s': %s",
|
||||||
|
self.name,
|
||||||
|
item,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
else:
|
||||||
|
value = item
|
||||||
|
runtime_tags = {}
|
||||||
|
|
||||||
|
if value is None:
|
||||||
|
logger.info(
|
||||||
|
"Skipping None value for '%s' in node '%s'",
|
||||||
|
msg_name,
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Merge base tags with runtime tags (runtime takes precedence)
|
||||||
|
tags = {**base_tags, **runtime_tags}
|
||||||
|
|
||||||
|
# Build the point
|
||||||
|
point = Point(measurement)
|
||||||
|
|
||||||
|
for tag_key, tag_value in tags.items():
|
||||||
|
point = point.tag(tag_key, str(tag_value))
|
||||||
|
|
||||||
|
point = point.field(field, value)
|
||||||
|
|
||||||
|
# Write the point
|
||||||
|
write_api.write(
|
||||||
|
bucket=self.bucket,
|
||||||
|
org=self.org,
|
||||||
|
record=point,
|
||||||
|
write_precision=precision,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Wrote to InfluxDB from node '%s': %s.%s=%s, tags=%s",
|
||||||
|
self.name,
|
||||||
|
measurement,
|
||||||
|
field,
|
||||||
|
value,
|
||||||
|
tags,
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("InfluxDB write error in node '%s': %s", self.name, e)
|
||||||
|
raise
|
||||||
|
|
||||||
|
def _query_data(self) -> dict:
|
||||||
|
"""
|
||||||
|
Query data from InfluxDB based on provides configuration.
|
||||||
|
|
||||||
|
:returns: Dict of port name to queried value.
|
||||||
|
:rtype: dict
|
||||||
|
"""
|
||||||
|
from influxdb_client import InfluxDBClient
|
||||||
|
|
||||||
|
results = {}
|
||||||
|
|
||||||
|
try:
|
||||||
|
with InfluxDBClient(url=self.url, token=self.token, org=self.org) as client:
|
||||||
|
query_api = client.query_api()
|
||||||
|
|
||||||
|
for spec in self.output_ports:
|
||||||
|
msg_name = spec.port
|
||||||
|
query_config = self.queries.get(msg_name, {})
|
||||||
|
|
||||||
|
measurement = query_config.get("measurement", msg_name)
|
||||||
|
field = query_config.get("field", "value")
|
||||||
|
tags = query_config.get("tags", {})
|
||||||
|
time_range = query_config.get("range", self.query_range)
|
||||||
|
aggregation = query_config.get("aggregation", "last")
|
||||||
|
|
||||||
|
# Build Flux query
|
||||||
|
flux_query = self._build_flux_query(
|
||||||
|
measurement=measurement,
|
||||||
|
field=field,
|
||||||
|
tags=tags,
|
||||||
|
time_range=time_range,
|
||||||
|
aggregation=aggregation,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Executing InfluxDB query for '%s' in node '%s': %s",
|
||||||
|
msg_name,
|
||||||
|
self.name,
|
||||||
|
flux_query,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Execute query
|
||||||
|
tables = query_api.query(flux_query, org=self.org)
|
||||||
|
|
||||||
|
# Extract result
|
||||||
|
value = self._extract_query_result(tables, spec)
|
||||||
|
|
||||||
|
if value is not None:
|
||||||
|
results[msg_name] = value
|
||||||
|
logger.info(
|
||||||
|
"Query result for '%s' in node '%s': %s",
|
||||||
|
msg_name,
|
||||||
|
self.name,
|
||||||
|
value,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.info(
|
||||||
|
"No data found for '%s' in node '%s'",
|
||||||
|
msg_name,
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("InfluxDB query error in node '%s': %s", self.name, e)
|
||||||
|
raise
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
def _build_flux_query(
|
||||||
|
self,
|
||||||
|
measurement: str,
|
||||||
|
field: str,
|
||||||
|
tags: dict,
|
||||||
|
time_range: str,
|
||||||
|
aggregation: str,
|
||||||
|
) -> str:
|
||||||
|
"""
|
||||||
|
Build a Flux query string.
|
||||||
|
|
||||||
|
:param measurement: Measurement name.
|
||||||
|
:type measurement: str
|
||||||
|
:param field: Field name.
|
||||||
|
:type field: str
|
||||||
|
:param tags: Tag filters.
|
||||||
|
:type tags: dict
|
||||||
|
:param time_range: Time range (e.g., "-1h").
|
||||||
|
:type time_range: str
|
||||||
|
:param aggregation: Aggregation function.
|
||||||
|
:type aggregation: str
|
||||||
|
:returns: Flux query string.
|
||||||
|
:rtype: str
|
||||||
|
"""
|
||||||
|
# Base query
|
||||||
|
query_parts = [
|
||||||
|
f'from(bucket: "{self.bucket}")',
|
||||||
|
f" |> range(start: {time_range})",
|
||||||
|
f' |> filter(fn: (r) => r["_measurement"] == "{measurement}")',
|
||||||
|
f' |> filter(fn: (r) => r["_field"] == "{field}")',
|
||||||
|
]
|
||||||
|
|
||||||
|
# Add tag filters
|
||||||
|
for tag_key, tag_value in tags.items():
|
||||||
|
query_parts.append(
|
||||||
|
f' |> filter(fn: (r) => r["{tag_key}"] == "{tag_value}")'
|
||||||
|
)
|
||||||
|
|
||||||
|
# Add aggregation
|
||||||
|
aggregation_map = {
|
||||||
|
"mean": "mean()",
|
||||||
|
"last": "last()",
|
||||||
|
"first": "first()",
|
||||||
|
"max": "max()",
|
||||||
|
"min": "min()",
|
||||||
|
"sum": "sum()",
|
||||||
|
"count": "count()",
|
||||||
|
}
|
||||||
|
|
||||||
|
if aggregation in aggregation_map:
|
||||||
|
query_parts.append(f" |> {aggregation_map[aggregation]}")
|
||||||
|
else:
|
||||||
|
# Default to last value
|
||||||
|
query_parts.append(" |> last()")
|
||||||
|
|
||||||
|
return "\n".join(query_parts)
|
||||||
|
|
||||||
|
def _extract_query_result(self, tables, spec: MessageSpec) -> Any:
|
||||||
|
"""
|
||||||
|
Extract a single value from query result tables.
|
||||||
|
|
||||||
|
:param tables: InfluxDB query result tables.
|
||||||
|
:param spec: The port the value is destined for.
|
||||||
|
:type spec: MessageSpec
|
||||||
|
:returns: Extracted and typed value, or None if no data.
|
||||||
|
:rtype: Any
|
||||||
|
"""
|
||||||
|
for table in tables:
|
||||||
|
for record in table.records:
|
||||||
|
value = record.get_value()
|
||||||
|
if value is not None:
|
||||||
|
try:
|
||||||
|
return spec.coerce(value)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
return value
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
def inject(self, outputs: dict | None = None) -> dict | None:
|
||||||
|
"""
|
||||||
|
Inject queried data into the pipeline.
|
||||||
|
|
||||||
|
For InfluxDbNode, inject performs a query operation and injects
|
||||||
|
the results into the pipeline. This is useful for trigger-style
|
||||||
|
usage where you want to periodically query InfluxDB.
|
||||||
|
|
||||||
|
:param outputs: Optional pre-set outputs (usually None for queries).
|
||||||
|
:type outputs: dict | None
|
||||||
|
:returns: Query results injected into the pipeline.
|
||||||
|
:rtype: dict | None
|
||||||
|
"""
|
||||||
|
if self._pipeline is None:
|
||||||
|
raise RuntimeError("Node must be bound to a pipeline to inject")
|
||||||
|
|
||||||
|
# Without given values, injecting means running the configured queries.
|
||||||
|
if not outputs:
|
||||||
|
if not self.provides:
|
||||||
|
return None
|
||||||
|
outputs = self._query_data()
|
||||||
|
|
||||||
|
return self._pipeline.trigger(self, self._to_messages(outputs))
|
||||||
@@ -0,0 +1,100 @@
|
|||||||
|
"""A small perceptron node, kept as a worked example of numeric logic."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import threading
|
||||||
|
from collections.abc import Iterable
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
from pydantic import BaseModel, ConfigDict
|
||||||
|
|
||||||
|
from app.flow.messages import MessageSpec
|
||||||
|
from app.flow.nodes.base import Node
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class MLPNode(Node):
|
||||||
|
"""
|
||||||
|
Multi-Layer Perceptron node for neural network processing in pipelines.
|
||||||
|
|
||||||
|
This node implements a simple single-layer neural network that applies
|
||||||
|
weights and biases to input values. Weights and biases are randomly
|
||||||
|
initialized using the provided random number generator.
|
||||||
|
|
||||||
|
The computation follows the standard neural network formula:
|
||||||
|
output = weights @ inputs + biases
|
||||||
|
|
||||||
|
:param requires: Input messages consumed by this node.
|
||||||
|
:type requires: MessageSpec | list[MessageSpec]
|
||||||
|
:param provides: Output messages produced by this node.
|
||||||
|
:type provides: MessageSpec | list[MessageSpec]
|
||||||
|
:param params: Parameters dict containing:
|
||||||
|
- ``rng`` (numpy.random.Generator): Random number generator for weight initialization
|
||||||
|
- Additional node parameters
|
||||||
|
:type params: dict
|
||||||
|
:param name: Name for this node.
|
||||||
|
:type name: str
|
||||||
|
|
||||||
|
:example:
|
||||||
|
>>> import numpy as np
|
||||||
|
>>> rng = np.random.default_rng(seed=42)
|
||||||
|
>>> mlp = MLPNode(
|
||||||
|
... requires=[MessageSpec(name="input1", dtype=DType.FLOAT), MessageSpec(name="input2", dtype=DType.FLOAT)],
|
||||||
|
... provides=[MessageSpec(name="output", dtype=DType.FLOAT)],
|
||||||
|
... params={"rng": rng},
|
||||||
|
... name="mlp_layer1",
|
||||||
|
... )
|
||||||
|
"""
|
||||||
|
|
||||||
|
class Params(BaseModel):
|
||||||
|
"""Weights are drawn from ``seed``, so a node reloads identically."""
|
||||||
|
|
||||||
|
model_config = ConfigDict(extra="allow")
|
||||||
|
|
||||||
|
seed: int = 0
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
requires: MessageSpec | Iterable[MessageSpec] = (),
|
||||||
|
provides: MessageSpec | Iterable[MessageSpec] = (),
|
||||||
|
params: dict[str, Any] | None = None,
|
||||||
|
name: str | None = None,
|
||||||
|
):
|
||||||
|
super().__init__(
|
||||||
|
self._forward,
|
||||||
|
requires=requires,
|
||||||
|
provides=provides,
|
||||||
|
params=params,
|
||||||
|
name=name or "mlp",
|
||||||
|
)
|
||||||
|
cfg = self.Params.model_validate(self.params)
|
||||||
|
|
||||||
|
rng = np.random.default_rng(seed=cfg.seed)
|
||||||
|
num_inputs = max(1, len(self.input_ports))
|
||||||
|
num_outputs = max(1, len(self.output_ports))
|
||||||
|
self.weights = rng.normal(loc=1, size=(num_outputs, num_inputs))
|
||||||
|
self.biases = rng.normal(loc=0, size=(num_outputs,))
|
||||||
|
|
||||||
|
def _forward(
|
||||||
|
self, params: dict[str, Any], **kwargs: float
|
||||||
|
) -> dict[str, Any] | None:
|
||||||
|
"""Apply ``weights @ inputs + biases`` to the incoming values."""
|
||||||
|
if not self.output_ports:
|
||||||
|
return None
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Executing MLP node in thread %s: %s",
|
||||||
|
threading.current_thread().name,
|
||||||
|
self.id,
|
||||||
|
)
|
||||||
|
|
||||||
|
if kwargs:
|
||||||
|
input_array = np.array([float(v) for v in kwargs.values()])
|
||||||
|
else:
|
||||||
|
input_array = np.array([1.0]) # Bias only, for source nodes.
|
||||||
|
|
||||||
|
outputs = np.dot(self.weights, input_array) + self.biases
|
||||||
|
return {p.port: float(outputs[i]) for i, p in enumerate(self.output_ports)}
|
||||||
@@ -0,0 +1,579 @@
|
|||||||
|
"""MQTT nodes: a subscriber that wakes the graph, a publisher that speaks for it."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
from collections.abc import Iterable
|
||||||
|
from enum import Enum
|
||||||
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
|
|
||||||
|
from app.flow.messages import MessageSpec
|
||||||
|
from app.flow.nodes.base import Node
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from fastapi import FastAPI
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# 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.
|
||||||
|
PUBLISH_QUEUE_SIZE = 256
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
: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)
|
||||||
|
|
||||||
|
__slots__ = (
|
||||||
|
"topics",
|
||||||
|
"mode",
|
||||||
|
"broker_host",
|
||||||
|
"broker_port",
|
||||||
|
"username",
|
||||||
|
"password",
|
||||||
|
"client_id",
|
||||||
|
"qos",
|
||||||
|
"retain",
|
||||||
|
"keepalive",
|
||||||
|
"_topic_to_ports",
|
||||||
|
"_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
|
||||||
|
|
||||||
|
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.
|
||||||
|
self._topic_to_ports: dict[str, list[str]] = {}
|
||||||
|
for port, topic in self.topics.items():
|
||||||
|
self._topic_to_ports.setdefault(topic, []).append(port)
|
||||||
|
|
||||||
|
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
|
||||||
|
|
||||||
|
# Runtime state
|
||||||
|
self._subscription_task: asyncio.Task | None = None
|
||||||
|
self._mqtt_client = None
|
||||||
|
self._stop_event: asyncio.Event | None = None
|
||||||
|
self._publish_queue: asyncio.Queue[dict] | None = None
|
||||||
|
self._publisher_task: asyncio.Task | 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, **kwargs) -> dict | 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, **kwargs) -> dict | 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], data: dict) -> 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,
|
||||||
|
) 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) -> 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,
|
||||||
|
) as client:
|
||||||
|
await self._publish_with(client, data)
|
||||||
|
|
||||||
|
async def _publish_with(self, client: Any, data: dict) -> 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
|
||||||
|
|
||||||
|
payload = 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=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:
|
||||||
|
self._publisher_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._publisher_task
|
||||||
|
except (asyncio.CancelledError, Exception): # noqa: B014 - shutting down
|
||||||
|
pass
|
||||||
|
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:
|
||||||
|
self._subscription_task.cancel()
|
||||||
|
try:
|
||||||
|
await self._subscription_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
self._subscription_task = None
|
||||||
|
self._stop_event = None
|
||||||
|
logger.info(
|
||||||
|
"Stopped MQTT subscription for node '%s'",
|
||||||
|
self.name,
|
||||||
|
)
|
||||||
|
|
||||||
|
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,
|
||||||
|
) 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._topic_to_ports.get(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 port as a key;
|
||||||
|
# anything else is the value itself.
|
||||||
|
if isinstance(parsed, dict) and port in parsed:
|
||||||
|
value = parsed[port]
|
||||||
|
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()
|
||||||
|
)
|
||||||
@@ -15,6 +15,7 @@ from app.core.config import settings
|
|||||||
from app.flow import logs
|
from app.flow import logs
|
||||||
from app.flow.controller import FlowController
|
from app.flow.controller import FlowController
|
||||||
from app.flow.events import event_bus
|
from app.flow.events import event_bus
|
||||||
|
from app.flow.nodes.http import close_shared_client
|
||||||
from app.flow.plugins import load_plugins
|
from app.flow.plugins import load_plugins
|
||||||
from app.flow.secrets import init_secrets
|
from app.flow.secrets import init_secrets
|
||||||
from app.flow.state import MemoryState, RedisState, StateBackend
|
from app.flow.state import MemoryState, RedisState, StateBackend
|
||||||
@@ -75,6 +76,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
|
|||||||
finally:
|
finally:
|
||||||
watchdog_task.cancel()
|
watchdog_task.cancel()
|
||||||
await controller.stop()
|
await controller.stop()
|
||||||
|
close_shared_client()
|
||||||
if settings.MCP_ENABLED:
|
if settings.MCP_ENABLED:
|
||||||
from app.mcp.http import aclose
|
from app.mcp.http import aclose
|
||||||
|
|
||||||
|
|||||||
@@ -47,11 +47,16 @@ build-backend = "hatchling.build"
|
|||||||
strict = true
|
strict = true
|
||||||
exclude = ["venv", ".venv", "alembic"]
|
exclude = ["venv", ".venv", "alembic"]
|
||||||
|
|
||||||
# app/flow/nodes.py still carries the prototype's node classes: the integration
|
# The integration nodes still carry the prototype's annotations, and croniter and
|
||||||
# nodes need annotations of their own, plus stubs for croniter and influxdb_client.
|
# influxdb_client ship no stubs. Shrinking one module at a time as each integration
|
||||||
# It is being revisited per integration; every module around it is checked strictly.
|
# is revisited; the rest of the package, and everything around it, is checked strictly.
|
||||||
[[tool.mypy.overrides]]
|
[[tool.mypy.overrides]]
|
||||||
module = ["app.flow.nodes"]
|
module = [
|
||||||
|
"app.flow.nodes.delay",
|
||||||
|
"app.flow.nodes.http",
|
||||||
|
"app.flow.nodes.influx",
|
||||||
|
"app.flow.nodes.mqtt",
|
||||||
|
]
|
||||||
ignore_errors = true
|
ignore_errors = true
|
||||||
|
|
||||||
[tool.ruff]
|
[tool.ruff]
|
||||||
|
|||||||
@@ -0,0 +1,66 @@
|
|||||||
|
"""The outbound nodes reuse one connection instead of opening one per message."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
from app.flow.messages import DType, MessageSpec
|
||||||
|
from app.flow.nodes import MqttNode
|
||||||
|
from app.flow.nodes.http import close_shared_client, shared_client
|
||||||
|
|
||||||
|
|
||||||
|
def test_http_senders_share_one_pooled_client():
|
||||||
|
first = shared_client()
|
||||||
|
try:
|
||||||
|
assert shared_client() is first
|
||||||
|
finally:
|
||||||
|
close_shared_client()
|
||||||
|
|
||||||
|
# Closing lets the next request build a fresh one rather than reusing a
|
||||||
|
# closed pool.
|
||||||
|
assert shared_client() is not first
|
||||||
|
close_shared_client()
|
||||||
|
|
||||||
|
|
||||||
|
def _publisher() -> MqttNode:
|
||||||
|
node = MqttNode(
|
||||||
|
requires=[MessageSpec(name="setpoint", port="setpoint", dtype=DType.FLOAT)],
|
||||||
|
params={"topic": {"setpoint": "heating/setpoint"}},
|
||||||
|
)
|
||||||
|
node.assign_flow("heating", "out")
|
||||||
|
return node
|
||||||
|
|
||||||
|
|
||||||
|
def test_a_started_publisher_queues_instead_of_connecting():
|
||||||
|
"""The handler runs on a worker thread; it must not block on the broker."""
|
||||||
|
node = _publisher()
|
||||||
|
|
||||||
|
async def scenario() -> None:
|
||||||
|
node._publish_queue = asyncio.Queue(maxsize=4)
|
||||||
|
node._loop = asyncio.get_running_loop()
|
||||||
|
|
||||||
|
await asyncio.to_thread(node._publisher_handler, {}, setpoint=21.0)
|
||||||
|
# call_soon_threadsafe lands on the next loop pass.
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert node._publish_queue.qsize() == 1
|
||||||
|
assert node._publish_queue.get_nowait() == {"setpoint": 21.0}
|
||||||
|
|
||||||
|
asyncio.run(scenario())
|
||||||
|
|
||||||
|
|
||||||
|
def test_a_full_publish_queue_drops_the_oldest():
|
||||||
|
"""A broker that cannot keep up must not grow the queue without bound."""
|
||||||
|
node = _publisher()
|
||||||
|
health: list[tuple[str, str | None]] = []
|
||||||
|
node._on_health = lambda _n, status, detail: health.append((status, detail))
|
||||||
|
|
||||||
|
async def scenario() -> None:
|
||||||
|
queue: asyncio.Queue[dict] = asyncio.Queue(maxsize=2)
|
||||||
|
for value in (1.0, 2.0, 3.0):
|
||||||
|
node._enqueue(queue, {"setpoint": value})
|
||||||
|
|
||||||
|
assert queue.qsize() == 2
|
||||||
|
assert queue.get_nowait() == {"setpoint": 2.0}
|
||||||
|
assert queue.get_nowait() == {"setpoint": 3.0}
|
||||||
|
assert health == [("degraded", "publish queue full")]
|
||||||
|
|
||||||
|
asyncio.run(scenario())
|
||||||
Reference in New Issue
Block a user