refactor: callbacks for csi data

This commit is contained in:
Christos Falas 2025-01-23 09:18:26 +00:00
parent 66e041cf73
commit 8b6ee4d137
No known key found for this signature in database
9 changed files with 153 additions and 61 deletions

View File

@ -11,6 +11,7 @@ pytest = "*"
matplotlib = "*"
scipy = "*"
scipy-stubs = "*"
typer = "*"
[dev-packages]

59
Pipfile.lock generated
View File

@ -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",

View File

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

36
src/cli.py Normal file
View File

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

View File

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

View File

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

View File

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

View File

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