diff --git a/where_fi/visualise/server/__init__.py b/where_fi/visualise/server/__init__.py new file mode 100644 index 0000000..d428bab --- /dev/null +++ b/where_fi/visualise/server/__init__.py @@ -0,0 +1,127 @@ +import logging +import multiprocessing as mp +import queue +import threading +import uuid +from concurrent import futures +from dataclasses import dataclass +from enum import Enum +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 .generated.figure_type import heatmap_pb2, line_pb2 + +logger = logging.getLogger(__name__) + +clients: dict[str, list[queue.Queue[figure_pb2.FigureData]]] = {} +clients_lock = threading.Lock() + + +class DataType(Enum): + RAW_CSI = 1 + PROCESSED_CSI = 2 + HEATMAP = 3 + + +@dataclass +class VisualiserData: + data: npt.NDArray[Any] + dtype: DataType + + +figures = { + "raw_phase": figure_pb2.Figure( + uuid=str(uuid.uuid4()), + title="Raw CSI Phase", + x_label="Subcarrier", + y_label="Phase", + ), + "processed_phase": figure_pb2.Figure( + uuid=str(uuid.uuid4()), + title="Preprocessed CSI Phase", + x_label="Subcarrier", + y_label="Phase", + ), + "aoa_heatmap": figure_pb2.Figure( + uuid=str(uuid.uuid4()), + title="Preprocessed CSI Phase", + x_label="Subcarrier", + y_label="Phase", + ), +} + + +class FigureServer(figure_pb2_grpc.FigureServiceServicer): + def GetFigure( + self, request: figure_pb2.FigureRequest, context: grpc.ServicerContext + ) -> Generator[figure_pb2.Figure, None, None]: + for figure in figures.values(): + 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: DataType, new_data: npt.NDArray[np.complex128]) -> None: + match dtype: + case DataType.RAW_CSI | DataType.PROCESSED_CSI: + uuid = ( + figures["raw_phase"].uuid + if dtype == DataType.RAW_CSI + else figures["processed_phase"].uuid + ) + lines = [ + line_pb2.LineChartData.Line( + y=np.angle(new_data)[:, i, 0], label=f"Antenna {i}" + ) + for i in range(new_data.shape[1]) + ] + linechart = line_pb2.LineChartData(lines=lines) + for q in clients.get(uuid, []): + q.put(figure_pb2.FigureData(uuid=uuid, line=linechart)) + case DataType.HEATMAP: + pass + + +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: + logger.info("Visualisation server shut down") + + 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() + server.wait_for_termination()