import logging import multiprocessing as mp import queue import threading import time 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 ..generated import figure_pb2, figure_pb2_grpc from . import figures @dataclass class VisualiserData: data: npt.NDArray[Any] dtype: figures.FigureId class FigureServer(figure_pb2_grpc.FigureServiceServicer): def __init__(self) -> None: self.logger = logging.getLogger(__name__) self.clients: dict[str, list[queue.Queue[figure_pb2.FigureData]]] = {} self.clients_lock = threading.Lock() def GetFigure( self, request: figure_pb2.FigureRequest, context: grpc.ServicerContext ) -> Generator[figure_pb2.Figure, None, None]: for figure_group in figures.all_figures.values(): for figure in figure_group.figures: yield figure def GetFigureUpdate( self, request: figure_pb2.FigureDataRequest, context: grpc.ServicerContext ) -> Generator[figure_pb2.FigureData, None, None]: self.logger.info( f"Received request for figure data stream for figure {request.uuid}" ) with self.clients_lock: q: queue.Queue[figure_pb2.FigureData] = queue.Queue() if request.uuid not in self.clients: self.clients[request.uuid] = [] self.clients[request.uuid].append(q) try: while context.is_active(): try: yield q.get(timeout=1) except queue.Empty: pass finally: with self.clients_lock: self.clients[request.uuid].remove(q) class Webapp: def __init__(self) -> None: self.logger = logging.getLogger(__name__) self.active = True self.figure_server = FigureServer() def add_data( self, dtype: figures.FigureId, new_data: npt.NDArray[np.complex128] ) -> None: self.logger.debug(f"Adding data to figure server {dtype}") if dtype not in figures.all_figures: self.logger.error(f"Figure {dtype} not found") return updates = figures.all_figures[dtype].update(new_data) for fig_id, update in updates.items(): for client in self.figure_server.clients.get(fig_id, []): client.put(update) def listen_for_data(self, data_queue: "mp.Queue[VisualiserData]") -> None: while self.active: self.logger.debug("Listening for data") try: data = data_queue.get(timeout=0.1) self.add_data(data.dtype, data.data) except queue.Empty: pass while not data_queue.empty(): data = data_queue.get() def start(self, data_queue: "mp.Queue[VisualiserData]") -> None: grpc_server = grpc.server(futures.ThreadPoolExecutor(max_workers=10)) figure_pb2_grpc.add_FigureServiceServicer_to_server( self.figure_server, grpc_server ) grpc_server.add_insecure_port("[::]:50051") grpc_server.add_insecure_port("0.0.0.0:50051") self.logger.info("Starting server on port 50051") grpc_server.start() self.logger.info("Server started") data_thread = threading.Thread(target=self.listen_for_data, args=(data_queue,)) data_thread.start() while self.active: time.sleep(1) self.logger.debug("Server is running") grpc_server.stop(0.5) data_thread.join()