Split off figure-serving from the data transformations needed to make the figures themselves (e.g. binning)
86 lines
2.6 KiB
Python
86 lines
2.6 KiB
Python
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 figures.all_figures[figure_group].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")
|