106 lines
3.4 KiB
Python
106 lines
3.4 KiB
Python
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:
|
|
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.5)
|
|
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()
|