Used to store the protobuffer messages into files, so that they can be plotted as part of the dissertation
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.Figure
|
|
|
|
|
|
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.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]:
|
|
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.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 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()
|