diff --git a/bec_ipython_client/bec_ipython_client/callbacks/move_device.py b/bec_ipython_client/bec_ipython_client/callbacks/move_device.py index 2f4249fb4..9f106f7c8 100644 --- a/bec_ipython_client/bec_ipython_client/callbacks/move_device.py +++ b/bec_ipython_client/bec_ipython_client/callbacks/move_device.py @@ -118,7 +118,7 @@ def get_device_values(self, force: bool = False) -> list: for dev in self.devices: val = self.data.get(dev) if val is None or force: - signal_data = self.device_manager.devices[dev].read(cached=True) + signal_data = self.device_manager.devices[dev].read(cached=not force) else: signal_data = val.signals # pylint: disable=protected-access diff --git a/bec_lib/bec_lib/bec_service.py b/bec_lib/bec_lib/bec_service.py index 7381501ae..678704216 100644 --- a/bec_lib/bec_lib/bec_service.py +++ b/bec_lib/bec_lib/bec_service.py @@ -283,6 +283,7 @@ def _send_service_status(self): info=messages.ServiceInfo(user=self._user, hostname=self._hostname), ), expire=6, + buffer=True, ) @property @@ -342,7 +343,7 @@ def _get_metrics(self): msg = messages.ServiceMetricMessage(name=self._service_name, metrics=data) try: self.connector.set_and_publish( - MessageEndpoints.metrics(self._service_name), msg, expire=30 + MessageEndpoints.metrics(self._service_name), msg, expire=30, buffer=True ) # pylint: disable=broad-except except Exception: diff --git a/bec_lib/bec_lib/logger.py b/bec_lib/bec_lib/logger.py index 96dace94a..6e24e5305 100644 --- a/bec_lib/bec_lib/logger.py +++ b/bec_lib/bec_lib/logger.py @@ -427,6 +427,8 @@ def _publish_log_message(self, msg: str | dict, service_name: str | None = None) ) }, max_size=10000, + buffer=True, + buffer_latest_only=False, ) return True except Exception: diff --git a/bec_lib/bec_lib/redis_connector/buffered_publisher.py b/bec_lib/bec_lib/redis_connector/buffered_publisher.py new file mode 100644 index 000000000..527903645 --- /dev/null +++ b/bec_lib/bec_lib/redis_connector/buffered_publisher.py @@ -0,0 +1,353 @@ +from __future__ import annotations + +import enum +import threading +import time +from collections import deque +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Callable, TypeAlias + +from bec_lib import messages +from bec_lib.logger import bec_logger + +if TYPE_CHECKING: # pragma: no cover + from redis.client import Pipeline + + from bec_lib.redis_connector import RedisConnector + +logger = bec_logger.logger + +RateLimitKey: TypeAlias = tuple[str, str] +MessageBuilder: TypeAlias = Callable[[], messages.BECMessage] +TopicMetrics: TypeAlias = dict[str, int] + + +class _PublishMethod(str, enum.Enum): + SET = "set" + SET_AND_PUBLISH = "set_and_publish" + SEND = "send" + XADD = "xadd" + + +@dataclass +class _PendingPublish: + """A queued publish request retained until the next shared flush.""" + + method: _PublishMethod + topic: str + msg: MessageBuilder | messages.BECMessage + kwargs: dict[str, Any] = field(default_factory=dict) + + +@dataclass +class _TopicState: + pending: deque[_PendingPublish] = field(default_factory=deque) + dropped_messages: int = 0 + + +class BufferedPublisher: + """Buffered Redis pipeline writes with immediate first publish per key.""" + + def __init__( + self, + connector: RedisConnector, + rate_limit_s: float = 0.1, + max_buffered_messages: int = 50, + log_metrics_period_s: float = 30, + ) -> None: + """ + Create a shared-cadence buffered publisher. + + Args: + connector (RedisConnector): Redis connector used for publishing. + rate_limit_s (float): Shared flush interval in seconds for pending updates. + + """ + self._connector = connector + self.rate_limit_s = rate_limit_s + self._max_buffered_messages = max_buffered_messages + self._lock = threading.Lock() + self._states: dict[RateLimitKey, _TopicState] = {} + self._ready: deque[_PendingPublish] = deque() + self._next_flush_at: float | None = None + self._evict_next_cycle: set[RateLimitKey] = set() + self._pending_event = threading.Event() + self._stop_event = threading.Event() + self._log_metrics_period_s = log_metrics_period_s + self._metrics_window_started_at = time.monotonic() + self._metrics_last_logged_at = self._metrics_window_started_at + self._metrics_sent_messages = 0 + self._metrics_flushes = 0 + self._metrics_min_batch_size: int | None = None + self._metrics_max_batch_size = 0 + self._metrics_dropped_messages = 0 + self._metrics_messages_by_topic: TopicMetrics = {} + self._thread = threading.Thread( + target=self._dispatch_pending, name="device-event-rate-limiter", daemon=True + ) + self._thread.start() + + def execute( + self, method: _PublishMethod | str, topic: str, builder: MessageBuilder, **kwargs: Any + ) -> None: + """ + Queue a buffered Redis operation. + + Args: + method (_PublishMethod | str): Redis operation to perform. + topic (str): Redis endpoint to update. + builder (MessageBuilder): Callback that builds the message at flush + time. + """ + normalized_method = _PublishMethod(method) + self._publish_rate_limited( + (normalized_method.value, topic), normalized_method, topic, builder, kwargs=kwargs + ) + + def shutdown(self) -> None: + """Stop the worker thread and flush any queued requests.""" + self._stop_event.set() + self._pending_event.set() + if self._thread.is_alive(): + self._thread.join(timeout=1) + + ################################################################################### + ################# Internal helper methods ######################################### + ################################################################################### + + def _dispatch_pending(self) -> None: + """ + Run the worker loop that flushes pending requests when they become due. + The loop is executed in a separate thread. + """ + while True: + pending_requests = self._collect_due_pending() + if pending_requests: + self._flush_requests(pending_requests) + continue + + if self._stop_event.is_set(): + return + + timeout = self._next_timeout() + self._pending_event.wait(timeout=timeout) + self._pending_event.clear() + + def _next_timeout(self) -> float | None: + """ + Return the wait time until the next shared flush. + + Returns: + float | None: Seconds until the next pending flush, or `None` if + nothing is pending. + """ + now = time.monotonic() + with self._lock: + if self._next_flush_at is None: + return None + return max(0.0, self._next_flush_at - now) + + def _collect_due_pending(self) -> list[_PendingPublish]: + """ + Collect requests that are ready to be flushed. + We also evict states that remained empty for a full flush cycle to + prevent unbounded memory growth. + + Returns: + list[_PendingPublish]: Ready-to-dispatch requests collected from the + immediate queue and the shared pending set. + """ + now = time.monotonic() + ready: list[_PendingPublish] = [] + with self._lock: + while self._ready: + ready.append(self._ready.popleft()) + should_flush_all_pending = self._stop_event.is_set() + if not should_flush_all_pending and ( + self._next_flush_at is None or self._next_flush_at > now + ): + return ready + + keys_to_evict = self._evict_next_cycle + self._evict_next_cycle = set() + flushed_keys: set[RateLimitKey] = set() + for key, state in self._states.items(): + if not state.pending: + continue + while state.pending: + ready.append(state.pending.popleft()) + flushed_keys.add(key) + self._next_flush_at = None + for key in keys_to_evict: + if key in self._states and not self._states[key].pending: + del self._states[key] + self._evict_next_cycle = flushed_keys + return ready + + def _dispatch_request(self, request: _PendingPublish, pipe: Pipeline) -> None: + """Queue a single request onto the Redis pipeline. + + Args: + request (_PendingPublish): Pending request to enqueue. + pipe (Pipeline): Redis pipeline that accumulates the write. + + Raises: + AttributeError: If the request's method is not recognized. + """ + msg = request.msg() if callable(request.msg) else request.msg + getattr(self._connector._managed_connection, request.method.value)( + request.topic, msg, pipe=pipe, **request.kwargs + ) + + def _flush_requests(self, requests: list[_PendingPublish] | None) -> None: + """Flush a batch of queued requests through a single Redis pipeline. + + Args: + requests (list[_PendingPublish] | None): Requests to flush. `None` + entries are ignored. + """ + if not requests: + return + + pipe = self._connector._managed_connection.pipeline() + for request in requests: + if request is None: + continue + try: + self._dispatch_request(request, pipe) + except Exception: + logger.exception("Failed to build or queue rate-limited device event callback") + + try: + pipe.execute() + self._record_flush_metrics(requests) + except Exception: + logger.exception("Failed to flush rate-limited device event pipeline") + + def _record_flush_metrics(self, requests: list[_PendingPublish]) -> None: + batch_size = len(requests) + if batch_size <= 0: + return + + topic_counts: TopicMetrics = {} + for request in requests: + topic_counts[request.topic] = topic_counts.get(request.topic, 0) + 1 + + with self._lock: + self._metrics_sent_messages += batch_size + self._metrics_flushes += 1 + for topic, count in topic_counts.items(): + self._metrics_messages_by_topic[topic] = ( + self._metrics_messages_by_topic.get(topic, 0) + count + ) + if self._metrics_min_batch_size is None: + self._metrics_min_batch_size = batch_size + else: + self._metrics_min_batch_size = min(self._metrics_min_batch_size, batch_size) + self._metrics_max_batch_size = max(self._metrics_max_batch_size, batch_size) + + now = time.monotonic() + if ( + self._log_metrics_period_s <= 0 + or now - self._metrics_last_logged_at < self._log_metrics_period_s + ): + return + + elapsed_s = max(now - self._metrics_window_started_at, 1e-9) + metrics_snapshot = { + "sent_messages": self._metrics_sent_messages, + "flushes": self._metrics_flushes, + "batch_min": self._metrics_min_batch_size or 0, + "batch_avg": self._metrics_sent_messages / self._metrics_flushes, + "batch_peak": self._metrics_max_batch_size, + "avg_message_rate_hz": self._metrics_sent_messages / elapsed_s, + "dropped_replaced": self._metrics_dropped_messages, + } + top_topics = sorted( + self._metrics_messages_by_topic.items(), key=lambda item: (-item[1], item[0]) + )[:5] + self._metrics_window_started_at = now + self._metrics_last_logged_at = now + self._metrics_sent_messages = 0 + self._metrics_flushes = 0 + self._metrics_min_batch_size = None + self._metrics_max_batch_size = 0 + self._metrics_dropped_messages = 0 + self._metrics_messages_by_topic = {} + + if metrics_snapshot["avg_message_rate_hz"] > 5: + # We don't need to log this if the message rate is low + top_topics_summary = ", ".join( + f"{topic}={count / elapsed_s:.1f}Hz ({count} msgs)" for topic, count in top_topics + ) + logger.info( + f"BufferedPublisher sent {metrics_snapshot['sent_messages']} messages across " + f"{metrics_snapshot['flushes']} flushes in {elapsed_s:.1f}s " + f"(batch min/avg/peak={metrics_snapshot['batch_min']}/" + f"{metrics_snapshot['batch_avg']:.1f}/{metrics_snapshot['batch_peak']}, " + f"avg msg rate={metrics_snapshot['avg_message_rate_hz']:.1f} Hz, " + f"dropped/replaced={metrics_snapshot['dropped_replaced']}, " + f"top topics: {top_topics_summary or 'n/a'})" + ) + try: + self._connector.publish_metrics("buffered_publisher", metrics_snapshot, separator="_") + except Exception: + logger.exception("Failed to publish buffered publisher metrics") + + def _publish_rate_limited( + self, + key: RateLimitKey, + method: _PublishMethod, + topic: str, + msg: MessageBuilder | messages.BECMessage, + kwargs: dict[str, Any] | None = None, + ) -> None: + """Queue a request for the next shared flush interval. + + Args: + key (RateLimitKey): Rate-limit bucket key derived from operation type and topic. + method (_PublishMethod): Redis operation to perform. + topic (str): Redis endpoint. + msg (MessageBuilder | messages.BECMessage): Message or callback that builds the message at flush time. + """ + now = time.monotonic() + request_kwargs = kwargs or {} + buffer_latest_only = request_kwargs.pop("buffer_latest_only", False) + pending_request = _PendingPublish( + method=method, topic=topic, msg=msg, kwargs=request_kwargs + ) + + with self._lock: + state = self._states.get(key) + if state is None: + state = _TopicState() + self._states[key] = state + self._evict_next_cycle.discard(key) + self._ready.append(pending_request) + else: + if buffer_latest_only: + if state.pending: + dropped = len(state.pending) + state.dropped_messages += dropped + self._metrics_dropped_messages += dropped + state.pending.clear() + state.pending.append(pending_request) + self._evict_next_cycle.discard(key) + if self._next_flush_at is None: + self._next_flush_at = now + self.rate_limit_s + self._pending_event.set() + return + + if len(state.pending) >= self._max_buffered_messages: + logger.warning( + f"Warning: Dropping message for {key[1]} due to exceeding max buffered messages ({self._max_buffered_messages})" + ) + state.dropped_messages += 1 + self._metrics_dropped_messages += 1 + return + + state.pending.append(pending_request) + self._evict_next_cycle.discard(key) + if self._next_flush_at is None: + self._next_flush_at = now + self.rate_limit_s + self._pending_event.set() diff --git a/bec_lib/bec_lib/redis_connector/hli.py b/bec_lib/bec_lib/redis_connector/hli.py index bcf6e4ecf..8b05b4788 100644 --- a/bec_lib/bec_lib/redis_connector/hli.py +++ b/bec_lib/bec_lib/redis_connector/hli.py @@ -4,7 +4,9 @@ from __future__ import annotations +import inspect import traceback +from functools import wraps from typing import Literal, Sequence from redis.client import Pipeline, Redis @@ -22,6 +24,7 @@ ) from bec_lib.messaging_hooks import MessagingEvent from bec_lib.messaging_services import NotificationMessageObject +from bec_lib.redis_connector.buffered_publisher import BufferedPublisher from bec_lib.serialization import MsgpackSerialization from .constants import IncompatibleMessageForEndpoint, IncompatibleRedisOperation, _BecMsgT @@ -31,6 +34,39 @@ logger = bec_logger.logger +def buffered_operation(func): + signature = inspect.signature(func) + parameter_names = list(signature.parameters) + topic_arg_name = parameter_names[1] + payload_arg_name = parameter_names[2] + + @wraps(func) + def wrapper(*args, **kwargs): + buffer = kwargs.pop("buffer", False) + buffer_latest_only = kwargs.pop("buffer_latest_only", True) + bound_args = signature.bind_partial(*args, **kwargs) + topic = bound_args.arguments[topic_arg_name] + pipe = bound_args.arguments.get("pipe") + buffered_publisher = getattr(args[0], "_buffered_publisher", None) + + if not buffer or pipe is not None or buffered_publisher is None: + return func(*args, **kwargs) + buffered_kwargs = dict(bound_args.arguments) + buffered_kwargs.pop("self", None) + buffered_kwargs.pop(topic_arg_name, None) + payload = buffered_kwargs.pop(payload_arg_name) + buffered_publisher.execute( + func.__name__, + topic, + lambda: payload, + buffer_latest_only=buffer_latest_only, + **buffered_kwargs, + ) + return None + + return wrapper + + class RedisConnector: """ Redis connector class. This class is a wrapper around the redis library providing @@ -44,6 +80,7 @@ def __init__( bootstrap: list[str] | str, redis_cls: type[Redis] = Redis, name: str = "RedisConnector", + buffered_publisher_enabled: bool = True, **kwargs, ): """ @@ -56,6 +93,7 @@ def __init__( **kwargs: additional keyword arguments to pass to the redis client. """ self._managed_connection = self.connector_cls(bootstrap, redis_cls, name, **kwargs) + self._buffered_publisher = BufferedPublisher(self) if buffered_publisher_enabled else None ################################## # SETUP AND CONFIG METHODS # @@ -86,6 +124,8 @@ def shutdown(self, per_thread_timeout_s: float | None = None): """ Shutdown the connector """ + if self._buffered_publisher is not None: + self._buffered_publisher.shutdown() return self._managed_connection.shutdown(per_thread_timeout_s) def register( @@ -245,7 +285,17 @@ def get_last(self, topic, key=None, count=1): return self._managed_connection.get_last(topic, key, count) @validate_endpoint("topic") - def set_and_publish(self, topic, msg, pipe=None, expire=None): + @buffered_operation + def set_and_publish( + self, + topic, + msg, + pipe=None, + expire=None, + *, + buffer: bool = False, + buffer_latest_only: bool = True, + ): return self._managed_connection.set_and_publish(topic, msg, pipe, expire) ############################## @@ -256,7 +306,16 @@ def raw_send(self, topic: str, msg, pipe=None): return self._managed_connection.raw_send(topic, msg, pipe) @validate_endpoint("topic") - def send(self, topic: str, msg: str | BECMessage, pipe: Pipeline | None = None) -> None: + @buffered_operation + def send( + self, + topic: str, + msg: str | BECMessage, + pipe: Pipeline | None = None, + *, + buffer: bool = False, + buffer_latest_only: bool = True, + ) -> None: return self._managed_connection.send(topic, msg, pipe) @validate_endpoint("topic") @@ -303,7 +362,19 @@ def mget(self, topics: list[str], pipe=None): return self._managed_connection.mget(topics, pipe) @validate_endpoint("topic") - def xadd(self, topic, msg_dict, max_size=None, pipe=None, expire=None, approximate=True): + @buffered_operation + def xadd( + self, + topic, + msg_dict, + max_size=None, + pipe=None, + expire=None, + approximate=True, + *, + buffer: bool = False, + buffer_latest_only: bool = True, + ): return self._managed_connection.xadd( topic, msg_dict, max_size=max_size, pipe=pipe, expire=expire, approximate=approximate ) diff --git a/bec_lib/bec_lib/redis_connector/validation.py b/bec_lib/bec_lib/redis_connector/validation.py index 8af4f5d27..89e9b5772 100644 --- a/bec_lib/bec_lib/redis_connector/validation.py +++ b/bec_lib/bec_lib/redis_connector/validation.py @@ -80,9 +80,10 @@ def validate_endpoint(endpoint_arg_name: str): def decorator( func: Callable[Concatenate[Any, str, P], Any], ) -> Callable[Concatenate[Any, EndpointInfo, P], Any]: - argspec = inspect.getfullargspec(func) + signature = inspect.signature(func) try: - argument_index = argspec.args.index(endpoint_arg_name) + parameter_names = list(signature.parameters) + argument_index = parameter_names.index(endpoint_arg_name) if argument_index != 1: raise ValueError except ValueError as e: diff --git a/bec_lib/bec_lib/tests/utils.py b/bec_lib/bec_lib/tests/utils.py index 19d3c9667..cbe82f29f 100644 --- a/bec_lib/bec_lib/tests/utils.py +++ b/bec_lib/bec_lib/tests/utils.py @@ -518,7 +518,7 @@ def __init__( bootstrap_server = bootstrap_server[0] if ":" not in bootstrap_server: bootstrap_server = f"{bootstrap_server}:0000" - super().__init__(bootstrap_server) + super().__init__(bootstrap_server, buffered_publisher_enabled=False) self.message_sent = [] self._get_buffer = {} self.store_data = store_data @@ -530,7 +530,7 @@ def log_error(self, *args, **kwargs): pass def shutdown(self, per_thread_timeout_s: float | None = None): - pass + super().shutdown(per_thread_timeout_s=per_thread_timeout_s) def register(self, *args, **kwargs): pass @@ -556,7 +556,7 @@ def raw_send(self, topic, msg, pipe=None): return self.message_sent.append({"queue": topic, "msg": msg}) - def send(self, topic, msg, pipe=None): + def send(self, topic, msg, pipe=None, *, buffer: bool = False, buffer_latest_only: bool = True): if not isinstance(msg, messages.BECMessage): raise TypeError("Message must be a BECMessage") return self.raw_send(topic, msg, pipe) @@ -572,7 +572,16 @@ def notify(self, event, message: str | NotificationMessageObject, pipe=None): pipe=pipe, ) - def set_and_publish(self, topic, msg, pipe=None, expire: int = None): + def set_and_publish( + self, + topic, + msg, + pipe=None, + expire: int = None, + *, + buffer: bool = False, + buffer_latest_only: bool = True, + ): if pipe: pipe._pipe_buffer.append(("set_and_publish", (topic.endpoint, msg), {"expire": expire})) return @@ -623,7 +632,17 @@ def lset(self, topic: str, index: int, msgs: str, pipe=None) -> None: def execute_pipeline(self, pipeline): pipeline.execute() - def xadd(self, topic, msg_dict, max_size=None, pipe=None, expire: int = None): + def xadd( + self, + topic, + msg_dict, + max_size=None, + pipe=None, + expire: int = None, + *, + buffer: bool = False, + buffer_latest_only: bool = True, + ): if pipe: pipe._pipe_buffer.append(("xadd", (topic, msg_dict), {"expire": expire})) return diff --git a/bec_lib/tests/conftest.py b/bec_lib/tests/conftest.py index 06462e39a..3b4d33942 100644 --- a/bec_lib/tests/conftest.py +++ b/bec_lib/tests/conftest.py @@ -34,7 +34,9 @@ def fake_redis_server(host, port, **kwargs): @pytest.fixture def connected_connector(): - connector = RedisConnector("localhost:1", redis_cls=fake_redis_server) + connector = RedisConnector( + "localhost:1", redis_cls=fake_redis_server, buffered_publisher_enabled=False + ) connector._managed_connection.flushall() try: yield connector diff --git a/bec_lib/tests/test_bec_service.py b/bec_lib/tests/test_bec_service.py index 54138c4b0..42c6f925d 100644 --- a/bec_lib/tests/test_bec_service.py +++ b/bec_lib/tests/test_bec_service.py @@ -27,7 +27,9 @@ class MagicMockConnector(RedisConnector): def __init__(self, *args, **kwargs): - super().__init__(*args, redis_cls=mock.MagicMock, **kwargs) + super().__init__( + *args, redis_cls=mock.MagicMock, buffered_publisher_enabled=False, **kwargs + ) @contextlib.contextmanager diff --git a/bec_server/bec_server/device_server/device_server.py b/bec_server/bec_server/device_server/device_server.py index 52d5719e4..7bcf2e6d9 100644 --- a/bec_server/bec_server/device_server/device_server.py +++ b/bec_server/bec_server/device_server/device_server.py @@ -844,8 +844,8 @@ def _read_device(self, instr: messages.DeviceInstructionMessage, new_status=True def _read_and_update_devices(self, devices: list[str], metadata: dict) -> list: start = time.time() - pipe = self.connector.pipeline() signal_container = [] + pipe = self.connector.pipeline() devices = self.device_manager.get_device_order(devices) for dev in devices: device_root = dev.split(".")[0] @@ -861,12 +861,12 @@ def _read_and_update_devices(self, devices: list[str], metadata: dict) -> list: self.connector.set_and_publish( MessageEndpoints.device_read(device_root), messages.DeviceMessage(signals=signals, metadata=metadata), - pipe, + pipe=pipe, ) self.connector.set_and_publish( MessageEndpoints.device_readback(device_root), messages.DeviceMessage(signals=signals, metadata=metadata), - pipe, + pipe=pipe, ) pipe.execute() logger.trace( @@ -876,7 +876,6 @@ def _read_and_update_devices(self, devices: list[str], metadata: dict) -> list: def _read_config_and_update_devices(self, devices: list[str], metadata: dict) -> list: start = time.time() - pipe = self.connector.pipeline() signal_container = [] devices = self.device_manager.get_device_order(devices) for dev in devices: @@ -891,9 +890,8 @@ def _read_config_and_update_devices(self, devices: list[str], metadata: dict) -> self.connector.set_and_publish( MessageEndpoints.device_read_configuration(dev), messages.DeviceMessage(signals=signals, metadata=metadata), - pipe, + buffer=True, ) - pipe.execute() logger.trace( f"Elapsed time for reading and updating status info: {(time.time() - start) * 1000} ms" ) diff --git a/bec_server/bec_server/device_server/devices/devicemanager.py b/bec_server/bec_server/device_server/devices/devicemanager.py index 3a0b4c003..a6cadf037 100644 --- a/bec_server/bec_server/device_server/devices/devicemanager.py +++ b/bec_server/bec_server/device_server/devices/devicemanager.py @@ -752,9 +752,7 @@ def _obj_callback_limit_change(self, *_args, obj: OphydObject, **kwargs): "high": {"value": obj.root.high_limit_travel.get()}, } dev_msg = messages.DeviceMessage(signals=limits) - pipe = self.connector.pipeline() - self.connector.set_and_publish(MessageEndpoints.device_limits(name), dev_msg, pipe=pipe) - pipe.execute() + self.connector.set_and_publish(MessageEndpoints.device_limits(name), dev_msg, buffer=True) def _obj_callback_readback(self, *_args, obj: OphydObject, **kwargs): if not obj.connected: @@ -763,9 +761,7 @@ def _obj_callback_readback(self, *_args, obj: OphydObject, **kwargs): signals = obj.root.read() metadata = self.devices.get(obj.root.name).metadata dev_msg = messages.DeviceMessage(signals=signals, metadata=metadata) - pipe = self.connector.pipeline() - self.connector.set_and_publish(MessageEndpoints.device_readback(name), dev_msg, pipe) - pipe.execute() + self.connector.set_and_publish(MessageEndpoints.device_readback(name), dev_msg, buffer=True) def _obj_callback_configuration(self, *_args, obj: OphydObject, **kwargs): if not obj.connected: @@ -777,11 +773,9 @@ def _obj_callback_configuration(self, *_args, obj: OphydObject, **kwargs): signals = obj.root.read_configuration() metadata = self.devices.get(obj.root.name).metadata dev_msg = messages.DeviceMessage(signals=signals, metadata=metadata) - pipe = self.connector.pipeline() self.connector.set_and_publish( - MessageEndpoints.device_read_configuration(name), dev_msg, pipe + MessageEndpoints.device_read_configuration(name), dev_msg, buffer=True ) - pipe.execute() @typechecked def _obj_callback_device_monitor_2d( diff --git a/bec_server/bec_server/scan_server/scan_server.py b/bec_server/bec_server/scan_server/scan_server.py index a3e79c7a5..eeee55ff3 100644 --- a/bec_server/bec_server/scan_server/scan_server.py +++ b/bec_server/bec_server/scan_server/scan_server.py @@ -127,3 +127,4 @@ def shutdown(self, per_thread_timeout_s: float | None = None) -> None: self.proc_manager.shutdown() self.device_manager.shutdown() self.queue_manager.shutdown() + super().shutdown(per_thread_timeout_s=per_thread_timeout_s) diff --git a/bec_server/tests/tests_device_server/test_device_server.py b/bec_server/tests/tests_device_server/test_device_server.py index 30501d1f8..a0a4265f0 100644 --- a/bec_server/tests/tests_device_server/test_device_server.py +++ b/bec_server/tests/tests_device_server/test_device_server.py @@ -872,12 +872,12 @@ def test_read_config_and_update_devices(device_server_mock, devices): res = [ msg for msg in device_server.connector.message_sent - if msg["queue"] == MessageEndpoints.device_read_configuration(device).endpoint + if msg["queue"] == MessageEndpoints.device_read_configuration(device) ] config = device_server.device_manager.devices[device].obj.read_configuration() msg = res[-1]["msg"] assert msg.content["signals"].keys() == config.keys() - assert res[-1]["queue"] == MessageEndpoints.device_read_configuration(device).endpoint + assert res[-1]["queue"] == MessageEndpoints.device_read_configuration(device) @pytest.mark.parametrize("device_manager_class", [DeviceManagerDS]) diff --git a/bec_server/tests/tests_scihub/conftest.py b/bec_server/tests/tests_scihub/conftest.py index 403514c18..5a1c1bebd 100644 --- a/bec_server/tests/tests_scihub/conftest.py +++ b/bec_server/tests/tests_scihub/conftest.py @@ -25,6 +25,9 @@ class SciHubMocked(SciHub): def _start_metrics_emitter(self): pass + def _start_update_service_info(self): + pass + def wait_for_service(self, name, status=BECStatus.RUNNING): pass @@ -82,7 +85,9 @@ def atlas_connector(connected_connector, connected_atlas_connector): }, ) scihub_mocked = SciHubMocked(config, ConnectorMock) + original_connector = scihub_mocked.connector scihub_mocked.connector = connected_connector # Replace with fakeredis connector + original_connector.shutdown() atlas_connector = AtlasConnector(scihub_mocked, connected_connector, connected_atlas_connector) with mock.patch.object(atlas_connector, "_load_environment"):