dissertation/where_fi/visualise/server/__init__.py
2025-05-16 00:41:31 +01:00

110 lines
3.6 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:
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()