Improve visualisation system #12

Merged
cfalas merged 3 commits from better-visualisation into main 2025-02-22 17:17:45 +02:00
Showing only changes of commit a94a9b4ec4 - Show all commits

View File

@ -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()