Measured with `make bench-engine` against a real Redis: 103.6 -> 164.4 messages a second on a five-node chain (p50 latency 2125 -> 1171 ms) and 34.8 -> 63.2 on a fan-out of twenty. Against the memory backend, which is what a pip install runs on, 262 -> 626. The two that bought most of it: - `StateBackend.record` puts a published value, its timestamp, its series and its version counter in one round trip. They were four calls building four pipelines, and a value crossing an edge pays them twice. A released rate-limit hold rides along instead of a DEL per port. - the readiness check reads a node's inputs and hands them to the node, rather than reading the triggering ones to count them and having the node read the same keys again a moment later. `apply_outputs` was a second copy of `_record_outputs` and is now the same code plus the event that distinguishes it. The rest, each small: - `_derive` builds a node-by-id map and a `consumes` index, so dispatching an item and publishing a value stop scanning every node in the installation. - `read_all` is memoised against the store revision — it sits on the publish path, so a dashboard slider was reading and validating every flow file per value. Same mechanism `_wiring` already uses. - the `message_value` source block is built once per node instead of per emission. - both timer threads ask the queue to promote only when something is actually due, which takes an idle engine from ~4 Redis round trips a second to one. - the shared httpx client is bounded (32 connections, one retry); its default pool is 100 with no per-host cap, so one slow endpoint could take it and every other sender node with it. - the MQTT and delay nodes no longer log a line per message at INFO. Robustness, in the same pass: - `MemoryWorkQueue._done` was a set nothing ever removed from — one entry per non-idempotent node per item, for the life of the process, in the default configuration. Capped, the way the Redis side expires its markers. - a saturated engine can claim from the due lane past the cascade limit. The capacity gate sits in front of the claim, so the due lane's priority — decided inside it — did not apply while every slot was held: a motor's stop was not behind the long nodes, it was unread. Only after a slot has genuinely failed to free for half a second, and briefly, so the backlog is not starved in turn. - `reclaim_stale` dispatches through that same gate. It could return sixty entries and push in-flight far past the limit the gate exists to hold. - a flow's nodes are stopped together rather than one after another. Each gets `NODE_STOP_TIMEOUT`, so a flow whose broker was unreachable took five seconds per node — long enough to outlast `REBUILD_WAIT` and 503 the deploy. - the worker pool and the HTTP client are closed on a thread, not on the event loop, and a run closes the state backend it built (on Redis, a client and a connection pool per run). - the five background tasks say something when they die. Each catches exceptions inside its loop, so one raised anywhere else left the engine serving with no metrics, no alerts or no artifact sweep, silently. `tests/flow/test_round_trips.py` counts the state operations one message costs — four, where it was about eleven — because none of the above would fail a behavioural test if it were undone. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01M6hPWS6YEbT1P8LxhhFb2T
794 lines
26 KiB
Python
794 lines
26 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."""
|
|
...
|
|
|
|
def close(self) -> None: # noqa: B027
|
|
"""Release whatever the backend holds outside this process.
|
|
|
|
A run builds a state backend of its own and drops it when it
|
|
finishes; on Redis that is a client and its connection pool, which
|
|
nothing was giving back. Concrete rather than abstract: a backend that
|
|
holds nothing has nothing to answer here.
|
|
"""
|
|
|
|
@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
|
|
"""
|
|
...
|
|
|
|
def record(
|
|
self,
|
|
values: dict[str, Any],
|
|
stamps: dict[str, float],
|
|
counters: list[str],
|
|
ts: float,
|
|
limits: dict[str, int] | None = None,
|
|
drop: list[str] | None = None,
|
|
) -> None:
|
|
"""Everything a published value owes state, in one round trip.
|
|
|
|
The value itself, the timestamp beside it, its series and its version
|
|
counter used to be four calls and four round trips; a backend that can
|
|
batch them does so here. This default is those four calls, so a
|
|
backend that cannot batch gains nothing and breaks nothing.
|
|
|
|
:param values: Message names mapped to the value just published.
|
|
:param stamps: The timestamp key of each of those messages.
|
|
:param counters: The version keys to bump.
|
|
:param ts: When they were published.
|
|
:param limits: How many points to keep per message.
|
|
:param drop: Keys to forget in the same trip — a released rate-limit
|
|
hold, which is written and cleared on this same path.
|
|
"""
|
|
for key in drop or ():
|
|
self.delete(key)
|
|
self.update({**values, **stamps})
|
|
self.append_history(values, ts, limits)
|
|
self.increment_multi(counters)
|
|
|
|
@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 record(
|
|
self,
|
|
values: dict[str, Any],
|
|
stamps: dict[str, float],
|
|
counters: list[str],
|
|
ts: float,
|
|
limits: dict[str, int] | None = None,
|
|
drop: list[str] | None = None,
|
|
) -> None:
|
|
"""The four writes under one acquisition of the one lock."""
|
|
with self._lock:
|
|
for key in drop or ():
|
|
self._data.pop(key, None)
|
|
self._history.pop(key, None)
|
|
self._data.update(values)
|
|
self._data.update(stamps)
|
|
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:
|
|
series = deque(series or (), maxlen=cap)
|
|
self._history[key] = series
|
|
series.append((ts, number))
|
|
for key in counters:
|
|
self._data[key] = self._data.get(key, 0) + 1
|
|
|
|
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 close(self) -> None:
|
|
"""Give the connection pool back."""
|
|
self._client.close()
|
|
|
|
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 record(
|
|
self,
|
|
values: dict[str, Any],
|
|
stamps: dict[str, float],
|
|
counters: list[str],
|
|
ts: float,
|
|
limits: dict[str, int] | None = None,
|
|
drop: list[str] | None = None,
|
|
) -> None:
|
|
"""Value, timestamp, series and version counter in one round trip.
|
|
|
|
These were four pipelines — four round trips — for one published
|
|
value, and a value crossing an edge pays them twice. Redis executes a
|
|
pipeline as one transaction, so batching them changes nothing about
|
|
what lands together and removes three of the four trips.
|
|
"""
|
|
pipe = self._client.pipeline()
|
|
for key in drop or ():
|
|
pipe.delete(self._key(key), self._history_key(key))
|
|
for key, value in {**values, **stamps}.items():
|
|
data = self._serialize(value)
|
|
if self._ttl:
|
|
pipe.setex(self._key(key), self._ttl, data)
|
|
else:
|
|
pipe.set(self._key(key), data)
|
|
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)
|
|
for key in counters:
|
|
pipe.incr(self._key(key))
|
|
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
|