A bool was excluded from the history as "not a measurement", so a true/false port had no curve in the node panel and none on an edge — only the word. It is recorded as 0/1 now and drawn as steps, since a bezier through two states slopes through readings that never happened. The axis is pinned to 0..1, so a flag that was never on sits at the floor rather than mid-box. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
683 lines
22 KiB
Python
683 lines
22 KiB
Python
"""
|
|
State management module with pluggable backends.
|
|
|
|
This module provides state storage backends for pipeline execution,
|
|
supporting both in-memory storage and Redis for distributed execution.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from abc import ABC, abstractmethod
|
|
from collections import deque
|
|
from collections.abc import Iterator
|
|
from contextlib import contextmanager
|
|
from threading import RLock
|
|
from typing import Any, cast
|
|
|
|
import orjson
|
|
import redis
|
|
|
|
#: A node may return a dict keyed by something other than a string, and the
|
|
#: stdlib encoder this replaced turned those keys into strings rather than
|
|
#: refusing them. Two differences remain and are both improvements: a
|
|
#: non-finite number is stored as ``null`` instead of the bare ``NaN`` that is
|
|
#: not JSON at all, and a datetime serialises rather than raising.
|
|
_JSON_OPTS = orjson.OPT_NON_STR_KEYS
|
|
|
|
# A sparkline only means something for a value that can be placed on an axis,
|
|
# so the history keeps those and nothing else. 120 points fill a panel-wide
|
|
# chart while leaving Redis a cache rather than a time-series database.
|
|
HISTORY_LIMIT = 120
|
|
|
|
|
|
def as_number(value: Any) -> float | None:
|
|
"""The plottable form of a value, or None if it cannot be placed on an axis."""
|
|
# A flag is not a measurement, but when it was on is worth seeing, and 0/1
|
|
# is what a step curve is drawn from. bool is an int subclass, so it would
|
|
# fall through the isinstance below either way.
|
|
if isinstance(value, bool):
|
|
return 1.0 if value else 0.0
|
|
if not isinstance(value, (int, float)):
|
|
return None
|
|
return float(value)
|
|
|
|
|
|
class StateBackend(ABC):
|
|
"""
|
|
Abstract base class for state storage backends.
|
|
|
|
Provides a dict-like interface for storing and retrieving state values
|
|
with support for distributed locking.
|
|
"""
|
|
|
|
@abstractmethod
|
|
def get(self, key: str, default: Any = None) -> Any:
|
|
"""
|
|
Get a value by key.
|
|
|
|
:param key: The key to retrieve.
|
|
:type key: str
|
|
:param default: Default value if key not found.
|
|
:type default: Any
|
|
:returns: The stored value or default.
|
|
:rtype: Any
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def set(self, key: str, value: Any) -> None:
|
|
"""
|
|
Set a value by key.
|
|
|
|
:param key: The key to set.
|
|
:type key: str
|
|
:param value: The value to store.
|
|
:type value: Any
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def exists(self, key: str) -> bool:
|
|
"""
|
|
Check if a key exists.
|
|
|
|
:param key: The key to check.
|
|
:type key: str
|
|
:returns: True if key exists.
|
|
:rtype: bool
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def delete(self, key: str) -> None:
|
|
"""Forget one key and any history kept for it."""
|
|
...
|
|
|
|
@abstractmethod
|
|
def clear(self) -> None:
|
|
"""Clear all keys in the state."""
|
|
...
|
|
|
|
@abstractmethod
|
|
def keys(self) -> list[str]:
|
|
"""
|
|
Get all keys.
|
|
|
|
:returns: List of all keys.
|
|
:rtype: list[str]
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def update(self, mapping: dict[str, Any]) -> None:
|
|
"""
|
|
Update multiple keys atomically.
|
|
|
|
:param mapping: Dict of key-value pairs to set.
|
|
:type mapping: dict[str, Any]
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
@contextmanager
|
|
def lock(self, timeout: float = 10.0) -> Iterator[None]:
|
|
"""
|
|
Acquire a lock for thread-safe operations.
|
|
|
|
:param timeout: Lock timeout in seconds.
|
|
:type timeout: float
|
|
:yields: Lock context.
|
|
"""
|
|
...
|
|
|
|
# -------------------------------------------------------------------------
|
|
# Atomic operations for synchronous node support
|
|
# -------------------------------------------------------------------------
|
|
|
|
@abstractmethod
|
|
def increment(self, key: str) -> int:
|
|
"""
|
|
Atomically increment a counter and return the new value.
|
|
|
|
:param key: The key to increment.
|
|
:type key: str
|
|
:returns: The new value after incrementing.
|
|
:rtype: int
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def get_multi(self, keys: list[str]) -> dict[str, Any]:
|
|
"""
|
|
Get multiple values atomically.
|
|
|
|
:param keys: List of keys to retrieve.
|
|
:type keys: list[str]
|
|
:returns: Dict mapping keys to values (missing keys have None).
|
|
:rtype: dict[str, Any]
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def get_present(self, keys: list[str]) -> dict[str, Any]:
|
|
"""The values of those keys that exist, in one round trip.
|
|
|
|
Unlike :meth:`get_multi`, a key that is absent is left out rather than
|
|
mapped to None — which is what reading a node's inputs needs, since a
|
|
message holding ``null`` is a message the node has, and one that was
|
|
never published is not.
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def increment_multi(self, keys: list[str]) -> None:
|
|
"""Bump several counters at once.
|
|
|
|
The new values are not returned: the only reader compares them against
|
|
what it saw last time and reads them back itself.
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def compare_and_swap_multi(
|
|
self,
|
|
expected: dict[str, Any],
|
|
updates: dict[str, Any],
|
|
) -> bool:
|
|
"""
|
|
Atomically check expected values and apply updates if all match.
|
|
|
|
This is the key primitive for race-condition-free synchronous node execution.
|
|
If any expected value doesn't match the current value, no updates are applied.
|
|
|
|
:param expected: Dict of keys to their expected values.
|
|
:type expected: dict[str, Any]
|
|
:param updates: Dict of keys to their new values (applied only if expected matches).
|
|
:type updates: dict[str, Any]
|
|
:returns: True if the swap succeeded, False if expected values didn't match.
|
|
:rtype: bool
|
|
"""
|
|
...
|
|
|
|
# -------------------------------------------------------------------------
|
|
# Message history
|
|
# -------------------------------------------------------------------------
|
|
|
|
@abstractmethod
|
|
def append_history(
|
|
self, values: dict[str, Any], ts: float, limits: dict[str, int] | None = None
|
|
) -> None:
|
|
"""
|
|
Record freshly published values in each message's capped history.
|
|
|
|
Non-numeric payloads are skipped — the history exists to be plotted.
|
|
|
|
:param values: Message names mapped to the value just published.
|
|
:type values: dict[str, Any]
|
|
:param ts: When they were published.
|
|
:param limits: How many points to keep per message, where a chart asks
|
|
for more than the default.
|
|
:type ts: float
|
|
"""
|
|
...
|
|
|
|
@abstractmethod
|
|
def history(self, key: str) -> list[tuple[float, float]]:
|
|
"""
|
|
The recorded ``(timestamp, value)`` pairs of one message, oldest first.
|
|
|
|
:param key: The message name.
|
|
:type key: str
|
|
:returns: At most ``HISTORY_LIMIT`` points; empty if nothing was recorded.
|
|
:rtype: list[tuple[float, float]]
|
|
"""
|
|
...
|
|
|
|
def __contains__(self, key: str) -> bool:
|
|
return self.exists(key)
|
|
|
|
def __getitem__(self, key: str) -> Any:
|
|
if not self.exists(key):
|
|
raise KeyError(key)
|
|
return self.get(key)
|
|
|
|
def __setitem__(self, key: str, value: Any) -> None:
|
|
self.set(key, value)
|
|
|
|
def __iter__(self) -> Iterator[str]:
|
|
return iter(self.keys())
|
|
|
|
def __len__(self) -> int:
|
|
return len(self.keys())
|
|
|
|
|
|
class MemoryState(StateBackend):
|
|
"""
|
|
In-memory state backend using a thread-safe dict.
|
|
|
|
Suitable for single-process execution and testing.
|
|
|
|
:example:
|
|
>>> state = MemoryState()
|
|
>>> state["key"] = "value"
|
|
>>> state["key"]
|
|
'value'
|
|
"""
|
|
|
|
__slots__ = ("_data", "_history", "_lock")
|
|
|
|
def __init__(self) -> None:
|
|
self._data: dict[str, Any] = {}
|
|
self._history: dict[str, deque[tuple[float, float]]] = {}
|
|
self._lock = RLock() # Reentrant lock for nested access
|
|
|
|
def get(self, key: str, default: Any = None) -> Any:
|
|
with self._lock:
|
|
return self._data.get(key, default)
|
|
|
|
def set(self, key: str, value: Any) -> None:
|
|
with self._lock:
|
|
self._data[key] = value
|
|
|
|
def exists(self, key: str) -> bool:
|
|
with self._lock:
|
|
return key in self._data
|
|
|
|
def delete(self, key: str) -> None:
|
|
with self._lock:
|
|
self._data.pop(key, None)
|
|
self._history.pop(key, None)
|
|
|
|
def clear(self) -> None:
|
|
with self._lock:
|
|
self._data.clear()
|
|
self._history.clear()
|
|
|
|
def keys(self) -> list[str]:
|
|
with self._lock:
|
|
return list(self._data.keys())
|
|
|
|
def update(self, mapping: dict[str, Any]) -> None:
|
|
with self._lock:
|
|
self._data.update(mapping)
|
|
|
|
@contextmanager
|
|
def lock(self, timeout: float = 10.0) -> Iterator[None]:
|
|
"""Acquire the internal lock."""
|
|
acquired = self._lock.acquire(timeout=timeout)
|
|
if not acquired:
|
|
raise TimeoutError("Failed to acquire lock")
|
|
try:
|
|
yield
|
|
finally:
|
|
self._lock.release()
|
|
|
|
def increment(self, key: str) -> int:
|
|
"""Atomically increment a counter and return the new value."""
|
|
with self._lock:
|
|
value: int = self._data.get(key, 0)
|
|
new_value = value + 1
|
|
self._data[key] = new_value
|
|
return new_value
|
|
|
|
def get_multi(self, keys: list[str]) -> dict[str, Any]:
|
|
"""Get multiple values atomically."""
|
|
with self._lock:
|
|
return {k: self._data.get(k) for k in keys}
|
|
|
|
def get_present(self, keys: list[str]) -> dict[str, Any]:
|
|
"""The values of those keys that are there."""
|
|
with self._lock:
|
|
return {k: self._data[k] for k in keys if k in self._data}
|
|
|
|
def increment_multi(self, keys: list[str]) -> None:
|
|
with self._lock:
|
|
for key in keys:
|
|
self._data[key] = self._data.get(key, 0) + 1
|
|
|
|
def compare_and_swap_multi(
|
|
self,
|
|
expected: dict[str, Any],
|
|
updates: dict[str, Any],
|
|
) -> bool:
|
|
"""Atomically check expected values and apply updates if all match."""
|
|
with self._lock:
|
|
# Check all expected values
|
|
for key, exp_value in expected.items():
|
|
actual = self._data.get(key)
|
|
if actual != exp_value:
|
|
return False
|
|
|
|
# All matched, apply updates
|
|
self._data.update(updates)
|
|
return True
|
|
|
|
def append_history(
|
|
self, values: dict[str, Any], ts: float, limits: dict[str, int] | None = None
|
|
) -> None:
|
|
"""Append the numeric values to their message's bounded series."""
|
|
with self._lock:
|
|
for key, value in values.items():
|
|
number = as_number(value)
|
|
if number is None:
|
|
continue
|
|
cap = (limits or {}).get(key, HISTORY_LIMIT)
|
|
series = self._history.get(key)
|
|
if series is None or series.maxlen != cap:
|
|
# A widget asking for a deeper series re-caps it in place.
|
|
series = deque(series or (), maxlen=cap)
|
|
self._history[key] = series
|
|
series.append((ts, number))
|
|
|
|
def history(self, key: str) -> list[tuple[float, float]]:
|
|
"""The recorded points of one message, oldest first."""
|
|
with self._lock:
|
|
return list(self._history.get(key, ()))
|
|
|
|
|
|
class RedisState(StateBackend):
|
|
"""
|
|
Redis-based state backend for distributed execution.
|
|
|
|
Values are stored as JSON — anything a node passes through must be
|
|
serializable. Uses a namespace prefix to isolate pipeline executions.
|
|
|
|
:param host: Redis host address.
|
|
:type host: str
|
|
:param port: Redis port number.
|
|
:type port: int
|
|
:param db: Redis database number.
|
|
:type db: int
|
|
:param namespace: Key prefix for isolation.
|
|
:type namespace: str
|
|
:param password: Redis password (optional).
|
|
:type password: str | None
|
|
:param ttl: Time-to-live for keys in seconds (optional).
|
|
:type ttl: int | None
|
|
|
|
:example:
|
|
>>> state = RedisState(host="localhost", namespace="my_pipeline")
|
|
>>> state["result"] = {"value": 42}
|
|
>>> state["result"]
|
|
{'value': 42}
|
|
"""
|
|
|
|
__slots__ = ("_client", "_namespace", "_ttl", "_lock_name")
|
|
|
|
def __init__(
|
|
self,
|
|
host: str = "localhost",
|
|
port: int = 6379,
|
|
db: int = 0,
|
|
namespace: str = "pipeline",
|
|
password: str | None = None,
|
|
ttl: int | None = None,
|
|
) -> None:
|
|
self._client = redis.Redis(
|
|
host=host,
|
|
port=port,
|
|
db=db,
|
|
password=password,
|
|
decode_responses=False,
|
|
)
|
|
self._namespace = namespace
|
|
self._ttl = ttl
|
|
self._lock_name = f"{namespace}:_lock"
|
|
|
|
def _key(self, key: str) -> str:
|
|
"""Generate namespaced key."""
|
|
return f"{self._namespace}:{key}"
|
|
|
|
def _serialize(self, value: Any) -> bytes:
|
|
"""Serialize value for storage."""
|
|
return orjson.dumps(value, option=_JSON_OPTS)
|
|
|
|
def _deserialize(self, data: bytes | None) -> Any:
|
|
"""Deserialize value from storage."""
|
|
return orjson.loads(data) if data else None
|
|
|
|
# redis-py types every command as a sync/async union; this is the
|
|
# synchronous client, so the results are narrowed where they are consumed.
|
|
def get(self, key: str, default: Any = None) -> Any:
|
|
data = cast(bytes | None, self._client.get(self._key(key)))
|
|
return self._deserialize(data) if data is not None else default
|
|
|
|
def set(self, key: str, value: Any) -> None:
|
|
data = self._serialize(value)
|
|
if self._ttl:
|
|
self._client.setex(self._key(key), self._ttl, data)
|
|
else:
|
|
self._client.set(self._key(key), data)
|
|
|
|
def exists(self, key: str) -> bool:
|
|
return bool(self._client.exists(self._key(key)))
|
|
|
|
def delete(self, key: str) -> None:
|
|
self._client.delete(self._key(key), self._history_key(key))
|
|
|
|
def clear(self) -> None:
|
|
"""Clear all keys in the namespace."""
|
|
pattern = f"{self._namespace}:*"
|
|
cursor = 0
|
|
while True:
|
|
cursor, keys = cast(
|
|
tuple[int, list[bytes]],
|
|
self._client.scan(cursor, match=pattern, count=100),
|
|
)
|
|
if keys:
|
|
self._client.delete(*keys)
|
|
if cursor == 0:
|
|
break
|
|
|
|
def keys(self) -> list[str]:
|
|
"""Get all keys in the namespace (without prefix)."""
|
|
pattern = f"{self._namespace}:*"
|
|
prefix_len = len(self._namespace) + 1
|
|
result = []
|
|
cursor = 0
|
|
while True:
|
|
cursor, keys = cast(
|
|
tuple[int, list[bytes]],
|
|
self._client.scan(cursor, match=pattern, count=100),
|
|
)
|
|
for k in keys:
|
|
key_str = k.decode("utf-8")
|
|
# Skip the lock key
|
|
if not key_str.endswith(":_lock"):
|
|
result.append(key_str[prefix_len:])
|
|
if cursor == 0:
|
|
break
|
|
return result
|
|
|
|
def update(self, mapping: dict[str, Any]) -> None:
|
|
"""Update multiple keys atomically using a pipeline."""
|
|
if not mapping:
|
|
return
|
|
pipe = self._client.pipeline()
|
|
for key, value in mapping.items():
|
|
data = self._serialize(value)
|
|
if self._ttl:
|
|
pipe.setex(self._key(key), self._ttl, data)
|
|
else:
|
|
pipe.set(self._key(key), data)
|
|
pipe.execute()
|
|
|
|
@contextmanager
|
|
def lock(self, timeout: float = 10.0) -> Iterator[None]:
|
|
"""
|
|
Acquire a distributed lock using Redis.
|
|
|
|
:param timeout: Lock timeout in seconds.
|
|
:type timeout: float
|
|
:yields: Lock context.
|
|
:raises TimeoutError: If lock cannot be acquired.
|
|
"""
|
|
redis_lock = self._client.lock(
|
|
self._lock_name,
|
|
timeout=timeout,
|
|
blocking_timeout=timeout,
|
|
)
|
|
acquired = redis_lock.acquire()
|
|
if not acquired:
|
|
raise TimeoutError("Failed to acquire Redis lock")
|
|
try:
|
|
yield
|
|
finally:
|
|
try:
|
|
redis_lock.release()
|
|
except redis.exceptions.LockError:
|
|
pass # Lock may have expired
|
|
|
|
def ping(self) -> bool:
|
|
"""
|
|
Check if Redis connection is available.
|
|
|
|
:returns: True if connected.
|
|
:rtype: bool
|
|
"""
|
|
try:
|
|
return bool(self._client.ping())
|
|
except redis.ConnectionError:
|
|
return False
|
|
|
|
def increment(self, key: str) -> int:
|
|
"""Atomically increment a counter using Redis INCR."""
|
|
return cast(int, self._client.incr(self._key(key)))
|
|
|
|
def increment_multi(self, keys: list[str]) -> None:
|
|
"""Bump every counter in one round trip rather than one INCR each."""
|
|
if not keys:
|
|
return
|
|
pipe = self._client.pipeline()
|
|
for key in keys:
|
|
pipe.incr(self._key(key))
|
|
pipe.execute()
|
|
|
|
def get_multi(self, keys: list[str]) -> dict[str, Any]:
|
|
"""Get multiple values atomically using Redis MGET."""
|
|
if not keys:
|
|
return {}
|
|
|
|
full_keys = [self._key(k) for k in keys]
|
|
values = cast(list[bytes | None], self._client.mget(full_keys))
|
|
|
|
result = {}
|
|
for key, value in zip(keys, values, strict=True):
|
|
result[key] = self._deserialize(value) if value is not None else None
|
|
|
|
return result
|
|
|
|
def get_present(self, keys: list[str]) -> dict[str, Any]:
|
|
"""The keys that are there, from one MGET.
|
|
|
|
A missing key comes back as ``None`` from Redis and a stored ``null``
|
|
comes back as the four bytes; only the raw reply tells the two apart,
|
|
which is why this is not ``get_multi`` with the Nones filtered out.
|
|
"""
|
|
if not keys:
|
|
return {}
|
|
values = cast(
|
|
list[bytes | None], self._client.mget([self._key(k) for k in keys])
|
|
)
|
|
return {
|
|
key: self._deserialize(raw)
|
|
for key, raw in zip(keys, values, strict=True)
|
|
if raw is not None
|
|
}
|
|
|
|
def __getitem__(self, key: str) -> Any:
|
|
# One GET rather than the base class's EXISTS-then-GET: a missing key
|
|
# and a stored `null` differ in the raw reply, so nothing has to ask
|
|
# twice.
|
|
raw = cast(bytes | None, self._client.get(self._key(key)))
|
|
if raw is None:
|
|
raise KeyError(key)
|
|
return self._deserialize(raw)
|
|
|
|
def compare_and_swap_multi(
|
|
self,
|
|
expected: dict[str, Any],
|
|
updates: dict[str, Any],
|
|
) -> bool:
|
|
"""
|
|
Atomically check expected values and apply updates using Redis WATCH/MULTI.
|
|
|
|
Uses optimistic locking with WATCH to detect concurrent modifications.
|
|
If any watched key changes before EXEC, the transaction is aborted.
|
|
|
|
:param expected: Dict of keys to their expected values.
|
|
:type expected: dict[str, Any]
|
|
:param updates: Dict of keys to their new values.
|
|
:type updates: dict[str, Any]
|
|
:returns: True if the swap succeeded, False otherwise.
|
|
:rtype: bool
|
|
"""
|
|
if not expected:
|
|
# No expectations, just set the updates
|
|
self.update(updates)
|
|
return True
|
|
|
|
watch_keys = [self._key(k) for k in expected.keys()]
|
|
|
|
try:
|
|
# Start watching the keys for changes
|
|
self._client.watch(*watch_keys)
|
|
|
|
# Check current values
|
|
current = self.get_multi(list(expected.keys()))
|
|
|
|
for key, exp_value in expected.items():
|
|
if current.get(key) != exp_value:
|
|
self._client.unwatch()
|
|
return False
|
|
|
|
# All values match - execute the transaction
|
|
pipe = self._client.pipeline(True) # True = use MULTI/EXEC
|
|
|
|
for key, value in updates.items():
|
|
data = self._serialize(value)
|
|
if self._ttl:
|
|
pipe.setex(self._key(key), self._ttl, data)
|
|
else:
|
|
pipe.set(self._key(key), data)
|
|
|
|
pipe.execute()
|
|
return True
|
|
|
|
except redis.WatchError:
|
|
# Another client modified one of the watched keys
|
|
return False
|
|
|
|
def _history_key(self, key: str) -> str:
|
|
return self._key(f"__history__:{key}")
|
|
|
|
def append_history(
|
|
self, values: dict[str, Any], ts: float, limits: dict[str, int] | None = None
|
|
) -> None:
|
|
"""Push the numeric values onto their capped list, in one round-trip."""
|
|
pipe = self._client.pipeline()
|
|
queued = False
|
|
for key, value in values.items():
|
|
number = as_number(value)
|
|
if number is None:
|
|
continue
|
|
history_key = self._history_key(key)
|
|
cap = (limits or {}).get(key, HISTORY_LIMIT)
|
|
pipe.lpush(history_key, orjson.dumps([ts, number]))
|
|
pipe.ltrim(history_key, 0, cap - 1)
|
|
if self._ttl:
|
|
pipe.expire(history_key, self._ttl)
|
|
queued = True
|
|
if queued:
|
|
pipe.execute()
|
|
|
|
def history(self, key: str) -> list[tuple[float, float]]:
|
|
"""The recorded points of one message, oldest first."""
|
|
entries = cast(list[bytes], self._client.lrange(self._history_key(key), 0, -1))
|
|
# LPUSH puts the newest first, a chart reads the other way round.
|
|
points: list[tuple[float, float]] = []
|
|
for entry in reversed(entries):
|
|
ts, value = orjson.loads(entry)
|
|
points.append((float(ts), float(value)))
|
|
return points
|