visualisation backend
This commit is contained in:
parent
964e379dbf
commit
a94a9b4ec4
127
where_fi/visualise/server/__init__.py
Normal file
127
where_fi/visualise/server/__init__.py
Normal 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()
|
||||||
Loading…
Reference in New Issue
Block a user