From ae2e95283b590516e7bc57c1395c34c81ea476dc Mon Sep 17 00:00:00 2001 From: Christos Falas Date: Mon, 30 Dec 2024 13:31:42 +0000 Subject: [PATCH] refactor: change Pipe to Queue for configurable maxsize --- src/__main__.py | 15 +++++++-------- src/aoa.py | 11 +++++++---- src/config.py | 2 -- src/ingest.py | 18 ++++++++++-------- src/visualise/__init__.py | 25 ++++++++++++------------- 5 files changed, 36 insertions(+), 35 deletions(-) diff --git a/src/__main__.py b/src/__main__.py index 7ffbd81..87b8f53 100644 --- a/src/__main__.py +++ b/src/__main__.py @@ -1,11 +1,11 @@ import logging import multiprocessing as mp -from multiprocessing.connection import Connection from typing import NamedTuple from . import ingest from . import config from . import visualise +from . import aoa logging.basicConfig( @@ -17,28 +17,27 @@ logging.basicConfig( class Receiver(NamedTuple): ip: str receiver: ingest.FeitReceiver - recv: Connection - send: Connection + queue: "mp.Queue[ingest.CSI]" receivers = [ - Receiver(ip, ingest.FeitReceiver(ip), *mp.Pipe(duplex=False)) + Receiver(ip, ingest.FeitReceiver(ip), mp.Queue(config.SAMPLE_RATE)) for ip in config.RECEIVE_IP_ADDRESS_LIST ] # Start injecting CSI frames transmitter = ingest.FeitTransmitter() -webapp_conns = mp.Pipe() +webapp_queue: "mp.Queue[aoa.AoA]" = mp.Queue(config.SAMPLE_RATE) -processor = ingest.CSIProcessor({r.ip: r.recv for r in receivers}, webapp_conns[0]) +processor = ingest.CSIProcessor({r.ip: r.queue for r in receivers}, webapp_queue) # Start webapp in background process -webapp = mp.Process(target=visualise.start, args=(webapp_conns[1],)) +webapp = mp.Process(target=visualise.start, args=(webapp_queue,)) webapp.start() receiver_processes = [ - mp.Process(target=r.receiver.listen, args=(r.send,)) for r in receivers + mp.Process(target=r.receiver.listen, args=(r.queue,)) for r in receivers ] for proc in receiver_processes: diff --git a/src/aoa.py b/src/aoa.py index 1d51392..600f90d 100644 --- a/src/aoa.py +++ b/src/aoa.py @@ -1,6 +1,7 @@ import numpy as np import numpy.typing as npt from . import config +from datetime import datetime import logging @@ -10,8 +11,9 @@ logger = logging.getLogger(__name__) class AoA: def __init__(self): self.historical_autocorr = np.array([]) - self.N_subcarriers = config.N_SUBCARRIERS - 2 - self.N_rx = len(config.ANTENNA_ORDER) + self.N_subcarriers = -1 + self.N_rx = -1 + self.timestamp = datetime.now() pass def smooth(self, data: npt.NDArray[np.complex128]): @@ -21,8 +23,8 @@ class AoA: N = data.shape[1] # Number of RX antennas T = data.shape[2] # Number of TX antennas - assert N == self.N_rx - assert M == self.N_subcarriers + self.N_subcarriers = M + self.N_rx = N logger.debug(f"Smoothing: Subcarriers: {M}, RX antennas: {N}, TX antennas: {T}") @@ -43,6 +45,7 @@ class AoA: return H_sm def update(self, data: npt.NDArray[np.complex128]): + self.timestamp = datetime.now() H_sm = self.smooth(data) auto_corr = np.matmul(H_sm, np.conj(H_sm).T) diff --git a/src/config.py b/src/config.py index 5ee7996..05614ca 100644 --- a/src/config.py +++ b/src/config.py @@ -31,8 +31,6 @@ ANTENNA_SPACING = 0.0285 # 2.85 cm CHANNEL_WIDTH = 20 FRAME_FORMAT = "HT" -N_SUBCARRIERS = 56 - CENTRAL_FREQUENCY_HZ = CENTRAL_FREQUENCY_MHZ * 1_000_000 C = 299_792_458 # m/s diff --git a/src/ingest.py b/src/ingest.py index 0600c1e..9e413b7 100644 --- a/src/ingest.py +++ b/src/ingest.py @@ -2,7 +2,6 @@ import logging import time import socket import multiprocessing as mp -from multiprocessing.connection import Connection import struct from datetime import datetime import numpy as np @@ -46,7 +45,7 @@ class FeitReceiver: f"{__name__}.{self.__class__.__name__}-{self.host}" ) - def listen(self, conn: Connection): + def listen(self, queue: "mp.Queue[CSI]"): prev_time = datetime.now() while True: # This is the max size of a UDP packet. The size of the actual CSI @@ -59,14 +58,16 @@ class FeitReceiver: f"Received CSI data after {datetime.now() - prev_time}" ) prev_time = datetime.now() - conn.send(csidata) + queue.put(csidata) except struct.error: self.logger.error("Failed to parse CSI data") class CSIProcessor: def __init__( - self, receiver_connections: dict[str, Connection], webserver: Connection + self, + receiver_connections: dict[str, "mp.Queue[CSI]"], + webserver: "mp.Queue[AoA]", ): self.pending_data: dict[str, tuple[datetime, CSI]] = {} self.pending_data_lock = mp.Lock() @@ -122,13 +123,14 @@ class CSIProcessor: processed = self.preprocess.preprocess(all_data) visualise.add_data(all_data, processed) self.aoa.update(processed) - self.webserver.send(self.aoa.E_n) + if not self.webserver.full(): + self.webserver.put(self.aoa) def process_forever(self): while True: - for ip, conn in self.connections.items(): - if conn.poll(): - self.add_data(ip, conn.recv()) + for ip, queue in self.connections.items(): + while not queue.empty(): + self.add_data(ip, queue.get()) if self.is_ready(): self.process_data() else: diff --git a/src/visualise/__init__.py b/src/visualise/__init__.py index 942e91a..1c27067 100644 --- a/src/visualise/__init__.py +++ b/src/visualise/__init__.py @@ -5,7 +5,7 @@ import numpy.typing as npt from simple_websocket import Server import time from datetime import datetime -from multiprocessing.connection import Connection +import multiprocessing as mp import matplotlib.pyplot as plt import io @@ -20,7 +20,7 @@ matplotlib.use("agg") app = Flask(__name__) sock = Sock(app) -aoa_conn = None +aoa_queue = None logger = logging.getLogger(__name__) @@ -92,7 +92,8 @@ def add_data( del subscriber_settings[subscriber] -def make_heatmap(max_tof: float): +def make_heatmap(aoa: AoA, max_tof: float): + logger.info(f"Making heatmap with aoa of {aoa.timestamp}") fig = plt.figure() ax = fig.add_axes([0, 0, 1, 1], polar=True) r = np.linspace(0, max_tof, 100) # Radius values @@ -112,20 +113,18 @@ def make_heatmap(max_tof: float): def gather_aoa(max_tof: float): - assert aoa_conn is not None + assert aoa_queue is not None prev_frame = datetime.now() while True: while (datetime.now() - prev_frame).total_seconds() < 1 / config.HEATMAP_FPS: time.sleep(0.01) - while aoa_conn.poll(): + while not aoa_queue.empty(): logger.debug("Receiving from aoa pipe") - E_n = aoa_conn.recv() - aoa.N_subcarriers = 54 - aoa.E_n = E_n + aoa = aoa_queue.get() prev_frame = datetime.now() - logger.debug("Generating heatmap") - buf = make_heatmap(max_tof) + logger.debug(f"Generating heatmap of time {aoa.timestamp}") + buf = make_heatmap(aoa, max_tof) yield (b"--frame\r\nContent-Type: image/jpeg\r\n\r\n" + buf.read() + b"\r\n") buf.close() @@ -142,8 +141,8 @@ def aoa_tof(): ) -def start(conn: Connection): - global aoa_conn +def start(conn: "mp.Queue[AoA]"): + global aoa_queue - aoa_conn = conn + aoa_queue = conn app.run(debug=True, use_reloader=False, host="0.0.0.0")