dissertation/where_fi/visualise/server/__init__.py
Christos Falas 8fe5995c5f
refactor figure visualisation
Split off figure-serving from the data transformations needed to make
the figures themselves (e.g. binning)
2025-03-04 17:30:26 +00:00

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")