diff --git a/Pipfile b/Pipfile index c4da6c2..121cf40 100644 --- a/Pipfile +++ b/Pipfile @@ -11,6 +11,7 @@ pytest = "*" matplotlib = "*" scipy = "*" scipy-stubs = "*" +typer = "*" [dev-packages] diff --git a/Pipfile.lock b/Pipfile.lock index 5e629e8..522566b 100644 --- a/Pipfile.lock +++ b/Pipfile.lock @@ -1,7 +1,7 @@ { "_meta": { "hash": { - "sha256": "c6ec3616be9da08134fd0a1eb5319ad8a9020abd1a929dd1196bbf5bf077f2e5" + "sha256": "dfb89f15ac6b380498f7a51a26ef1ba4b6e3a0cc01c8770dec4f2ef4b852371f" }, "pipfile-spec": 6, "requires": { @@ -292,6 +292,14 @@ "markers": "python_version >= '3.10'", "version": "==1.4.8" }, + "markdown-it-py": { + "hashes": [ + "sha256:355216845c60bd96232cd8d8c40e8f9765cc86f46880e43a8fd22dc1a1a8cab1", + "sha256:e3f60a94fa066dc52ec76661e37c851cb232d92f9886b15cb560aaada2df8feb" + ], + "markers": "python_version >= '3.8'", + "version": "==3.0.0" + }, "markupsafe": { "hashes": [ "sha256:0bff5e0ae4ef2e1ae4fdf2dfd5b76c75e5c2fa4132d05fc1b0dabcd20c7e28c4", @@ -400,6 +408,14 @@ "markers": "python_version >= '3.10'", "version": "==3.10.0" }, + "mdurl": { + "hashes": [ + "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", + "sha256:bb413d29f5eea38f31dd4754dd7377d4465116fb207585f97bf925588687c1ba" + ], + "markers": "python_version >= '3.7'", + "version": "==0.1.2" + }, "numpy": { "hashes": [ "sha256:059e6a747ae84fce488c3ee397cee7e5f905fd1bda5fb18c66bc41807ff119b2", @@ -567,6 +583,14 @@ "markers": "python_version >= '3.8'", "version": "==1.5.0" }, + "pygments": { + "hashes": [ + "sha256:61c16d2a8576dc0649d9f39e089b5f02bcd27fba10d8fb4dcc28173f7a45151f", + "sha256:9ea1544ad55cecf4b8242fab6dd35a93bbce657034b0611ee383099054ab6d8c" + ], + "markers": "python_version >= '3.8'", + "version": "==2.19.1" + }, "pyparsing": { "hashes": [ "sha256:93d9577b88da0bbea8cc8334ee8b918ed014968fd2ec383e868fb8afb1ccef84", @@ -592,6 +616,14 @@ "markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'", "version": "==2.9.0.post0" }, + "rich": { + "hashes": [ + "sha256:439594978a49a09530cff7ebc4b5c7103ef57baf48d5ea3184f21d9a2befa098", + "sha256:6049d5e6ec054bf2779ab3358186963bac2ea89175919d699e378b99738c2a90" + ], + "markers": "python_full_version >= '3.8.0'", + "version": "==13.9.4" + }, "scipy": { "hashes": [ "sha256:0c2f95de3b04e26f5f3ad5bb05e74ba7f68b837133a4492414b3afd79dfe540e", @@ -641,6 +673,14 @@ "markers": "python_version >= '3.10'", "version": "==1.14.1.6" }, + "shellingham": { + "hashes": [ + "sha256:7ecfff8f2fd72616f7481040475a65b2bf8af90a56c89140852d1120324e8686", + "sha256:8dbca0739d487e5bd35ab3ca4b36e11c4078f3a234bfce294b0a0291363404de" + ], + "markers": "python_version >= '3.7'", + "version": "==1.5.4" + }, "simple-websocket": { "hashes": [ "sha256:4af6069630a38ed6c561010f0e11a5bc0d4ca569b36306eb257cd9a192497c8c", @@ -657,6 +697,23 @@ "markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'", "version": "==1.17.0" }, + "typer": { + "hashes": [ + "sha256:7994fb7b8155b64d3402518560648446072864beefd44aa2dc36972a5972e847", + "sha256:a0588c0a7fa68a1978a069818657778f86abe6ff5ea6abf472f940a08bfe4f0a" + ], + "index": "pypi", + "markers": "python_version >= '3.7'", + "version": "==0.15.1" + }, + "typing-extensions": { + "hashes": [ + "sha256:04e5ca0351e0f3f85c6853954072df659d0d13fac324d0072316b67d7794700d", + "sha256:1a7ead55c7e559dd4dee8856e3a88b41225abfe1ce8df57b7c13915fe121ffb8" + ], + "markers": "python_version >= '3.8'", + "version": "==4.12.2" + }, "werkzeug": { "hashes": [ "sha256:54b78bf3716d19a65be4fceccc0d1d7b89e608834989dfae50ea87564639213e", diff --git a/src/__main__.py b/src/__main__.py index 1e258a2..71ea413 100644 --- a/src/__main__.py +++ b/src/__main__.py @@ -1,11 +1,6 @@ import logging -import multiprocessing as mp -from typing import NamedTuple -from . import ingest -from . import config -from . import visualise -from . import aoa +from . import cli logging.basicConfig( @@ -13,35 +8,5 @@ logging.basicConfig( format="%(asctime)s %(name)-40s %(levelname)-8s %(message)s", ) - -class Receiver(NamedTuple): - ip: ingest.Host - receiver: ingest.FeitReceiver - queue: "mp.Queue[ingest.CSI]" - - -receivers = [ - Receiver(ip, ingest.FeitReceiver(ip), mp.Queue(config.SAMPLE_RATE)) - for ip in config.RECEIVE_HOSTS -] - -# Start injecting CSI frames -transmitter = ingest.FeitTransmitter() - -webapp_queue: "mp.Queue[aoa.AoA]" = mp.Queue(config.SAMPLE_RATE) - -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_queue,)) -webapp.start() - -receiver_processes = [ - mp.Process(target=r.receiver.listen, args=(r.queue,)) for r in receivers -] - -for proc in receiver_processes: - proc.start() - -processing_thread = mp.Process(target=processor.process_forever) -processing_thread.start() +if __name__ == "__main__": + cli.app() diff --git a/src/cli.py b/src/cli.py new file mode 100644 index 0000000..40ae56e --- /dev/null +++ b/src/cli.py @@ -0,0 +1,36 @@ +import typer +import multiprocessing as mp + +from .collection import ingest +from .processing.preprocess import Preprocessor +from .processing.aoa import AoA +from . import config +from . import visualise + +import numpy as np +import numpy.typing as npt + +app = typer.Typer() + + +@app.command() +def heatmap(): + preprocessor = Preprocessor() + aoa = AoA() + webapp_queue: "mp.Queue[AoA]" = mp.Queue(config.SAMPLE_RATE) + # Start webapp in background process + webapp = mp.Process(target=visualise.start, args=(webapp_queue,)) + webapp.start() + + def callback(antenna_data: npt.NDArray[np.complex128]): + processed = preprocessor.preprocess(antenna_data) + # visualise.add_data(all_data, processed) + aoa.update(processed) + if not webapp_queue.full(): + webapp_queue.put(aoa) + + ingest.start_processing(callback) + + +if __name__ == "__main__": + app() diff --git a/src/csi.py b/src/collection/csi_frame.py similarity index 100% rename from src/csi.py rename to src/collection/csi_frame.py diff --git a/src/ingest.py b/src/collection/ingest.py similarity index 74% rename from src/ingest.py rename to src/collection/ingest.py index 579e27e..315aa73 100644 --- a/src/ingest.py +++ b/src/collection/ingest.py @@ -3,16 +3,18 @@ import time import socket import multiprocessing as mp import struct +from typing import Callable from datetime import datetime import numpy as np +import numpy.typing as npt +from typing import NamedTuple -from .csi import CSI -from .aoa import AoA -from .preprocess import Preprocessor -from . import config +from .csi_frame import CSI +from .. import config Host = tuple[str, int] +CSICallback = Callable[[npt.NDArray[np.complex128]], None] class FeitTransmitter: @@ -28,6 +30,8 @@ class FeitTransmitter: ) self.inject_server.send(b"stop\n") self.inject_server.send(inject_start_string.encode()) + self.logger = logging.getLogger(f"{__name__}.{self.__class__.__name__}") + self.logger.info(f"Connected to {config.INJECT_HOST}") class FeitReceiver: @@ -46,9 +50,11 @@ class FeitReceiver: self.logger = logging.getLogger( f"{__name__}.{self.__class__.__name__}-{self.host}" ) + self.logger.info(f"Connected to {self.host}") def listen(self, queue: "mp.Queue[CSI]"): prev_time = datetime.now() + self.logger.info("Listening for CSI data") while True: # This is the max size of a UDP packet. The size of the actual CSI # packet will depend on the frame format and channel width, which @@ -69,18 +75,12 @@ class CSIProcessor: def __init__( self, receiver_connections: dict[Host, "mp.Queue[CSI]"], - webserver: "mp.Queue[AoA]", ): self.pending_data: dict[Host, tuple[datetime, CSI]] = {} self.pending_data_lock = mp.Lock() self.logger = logging.getLogger(f"{__name__}.{self.__class__.__name__}") self.last_processed = datetime.now() - - self.preprocess = Preprocessor() - self.aoa = AoA() - self.connections = receiver_connections - self.webserver = webserver def add_data(self, host: Host, data: CSI): if ( @@ -111,7 +111,7 @@ class CSIProcessor: return False return True - def process_data(self): + def process_data(self, callback: CSICallback): self.last_processed = datetime.now() antenna_data = [ np.expand_dims(self.pending_data[ip][1].matrix[:, antenna], axis=2) @@ -122,19 +122,52 @@ class CSIProcessor: all_data = np.concat(antenna_data, axis=1) self.logger.info(f"Got final CSI data with shape {all_data.shape}") - processed = self.preprocess.preprocess(all_data) - visualise.add_data(all_data, processed) - self.aoa.update(processed) - if not self.webserver.full(): - self.webserver.put(self.aoa) + callback(all_data) - def process_forever(self): + def process_forever( + self, + callback: CSICallback, + ): while True: for ip, queue in self.connections.items(): while not queue.empty(): self.add_data(ip, queue.get()) if self.is_ready(): - self.process_data() + self.process_data(callback) else: self.logger.debug("Not all data is ready") - time.sleep(0.001) + time.sleep(0.0005) + + +class Receiver(NamedTuple): + ip: Host + receiver: FeitReceiver + queue: "mp.Queue[CSI]" + + +def start_processing(csi_callback: CSICallback): + receivers = [ + Receiver(ip, FeitReceiver(ip), mp.Queue(config.SAMPLE_RATE)) + for ip in config.RECEIVE_HOSTS + ] + + FeitTransmitter() + + # Start injecting CSI frames + processor = CSIProcessor({r.ip: r.queue for r in receivers}) + + receiver_processes = [ + mp.Process(target=r.receiver.listen, args=(r.queue,)) for r in receivers + ] + + for proc in receiver_processes: + proc.start() + + processing_thread = mp.Process( + target=processor.process_forever, args=(csi_callback,) + ) + processing_thread.start() + try: + processing_thread.join() + except KeyboardInterrupt: + return diff --git a/src/aoa.py b/src/processing/aoa.py similarity index 99% rename from src/aoa.py rename to src/processing/aoa.py index 600f90d..9f427d9 100644 --- a/src/aoa.py +++ b/src/processing/aoa.py @@ -1,6 +1,6 @@ import numpy as np import numpy.typing as npt -from . import config +from .. import config from datetime import datetime import logging diff --git a/src/preprocess.py b/src/processing/preprocess.py similarity index 99% rename from src/preprocess.py rename to src/processing/preprocess.py index 7cea971..f3a290c 100644 --- a/src/preprocess.py +++ b/src/processing/preprocess.py @@ -1,7 +1,7 @@ import numpy as np import logging from queue import Queue -from . import config +from .. import config import numpy.typing as npt from scipy.signal import correlate from scipy.signal import butter, sosfilt_zi, sosfilt diff --git a/src/visualise/__init__.py b/src/visualise/__init__.py index 1c27067..83f96c7 100644 --- a/src/visualise/__init__.py +++ b/src/visualise/__init__.py @@ -11,7 +11,7 @@ import matplotlib.pyplot as plt import io import logging -from ..aoa import AoA +from ..processing.aoa import AoA from .. import config import matplotlib