import logging import multiprocessing as mp import queue import threading from concurrent import futures from dataclasses import dataclass from typing import Any, Generator import grpc import numpy as np import numpy.typing as npt from . import figures from .generated import figure_pb2, figure_pb2_grpc logger = logging.getLogger(__name__) clients: dict[str, list[queue.Queue[figure_pb2.FigureData]]] = {} clients_lock = threading.Lock() @dataclass class VisualiserData: data: npt.NDArray[Any] dtype: figures.Figure class FigureServer(figure_pb2_grpc.FigureServiceServicer): def GetFigure( self, request: figure_pb2.FigureRequest, context: grpc.ServicerContext ) -> Generator[figure_pb2.Figure, None, None]: for figure_group in figures.Figure: for figure in figure_group.figure_class().figures: yield figure def GetFigureUpdate( self, request: figure_pb2.FigureDataRequest, context: grpc.ServicerContext ) -> Generator[figure_pb2.FigureData, None, None]: logger.info( f"Received request for figure data stream for figure {request.uuid}" ) with clients_lock: q: queue.Queue[figure_pb2.FigureData] = queue.Queue() if request.uuid not in clients: clients[request.uuid] = [] clients[request.uuid].append(q) try: while context.is_active(): try: yield q.get(timeout=1) except queue.Empty: pass finally: with clients_lock: clients[request.uuid].remove(q) def add_data(dtype: figures.Figure, new_data: npt.NDArray[np.complex128]) -> None: updates = figures.all_figures[dtype].update(new_data) for fig_id, update in updates.items(): for client in clients.get(fig_id, []): client.put(update) def listen_for_data(data_queue: "mp.Queue[VisualiserData]") -> None: while True: data = data_queue.get() add_data(data.dtype, data.data) def start(data_queue: "mp.Queue[VisualiserData]") -> None: server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) figure_pb2_grpc.add_FigureServiceServicer_to_server(FigureServer(), server) server.add_insecure_port("[::]:50051") server.add_insecure_port("0.0.0.0:50051") logger.info("Starting server on port 50051") server.start() logger.info("Server started") threading.Thread(target=listen_for_data, args=(data_queue,), daemon=True).start() try: server.wait_for_termination() except KeyboardInterrupt: logger.info("Exiting visualisation server")