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 = "*" matplotlib = "*"
scipy = "*" scipy = "*"
scipy-stubs = "*" scipy-stubs = "*"
typer = "*"
[dev-packages] [dev-packages]

59
Pipfile.lock generated
View File

@ -1,7 +1,7 @@
{ {
"_meta": { "_meta": {
"hash": { "hash": {
"sha256": "c6ec3616be9da08134fd0a1eb5319ad8a9020abd1a929dd1196bbf5bf077f2e5" "sha256": "dfb89f15ac6b380498f7a51a26ef1ba4b6e3a0cc01c8770dec4f2ef4b852371f"
}, },
"pipfile-spec": 6, "pipfile-spec": 6,
"requires": { "requires": {
@ -292,6 +292,14 @@
"markers": "python_version >= '3.10'", "markers": "python_version >= '3.10'",
"version": "==1.4.8" "version": "==1.4.8"
}, },
"markdown-it-py": {
"hashes": [
"sha256:355216845c60bd96232cd8d8c40e8f9765cc86f46880e43a8fd22dc1a1a8cab1",
"sha256:e3f60a94fa066dc52ec76661e37c851cb232d92f9886b15cb560aaada2df8feb"
],
"markers": "python_version >= '3.8'",
"version": "==3.0.0"
},
"markupsafe": { "markupsafe": {
"hashes": [ "hashes": [
"sha256:0bff5e0ae4ef2e1ae4fdf2dfd5b76c75e5c2fa4132d05fc1b0dabcd20c7e28c4", "sha256:0bff5e0ae4ef2e1ae4fdf2dfd5b76c75e5c2fa4132d05fc1b0dabcd20c7e28c4",
@ -400,6 +408,14 @@
"markers": "python_version >= '3.10'", "markers": "python_version >= '3.10'",
"version": "==3.10.0" "version": "==3.10.0"
}, },
"mdurl": {
"hashes": [
"sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8",
"sha256:bb413d29f5eea38f31dd4754dd7377d4465116fb207585f97bf925588687c1ba"
],
"markers": "python_version >= '3.7'",
"version": "==0.1.2"
},
"numpy": { "numpy": {
"hashes": [ "hashes": [
"sha256:059e6a747ae84fce488c3ee397cee7e5f905fd1bda5fb18c66bc41807ff119b2", "sha256:059e6a747ae84fce488c3ee397cee7e5f905fd1bda5fb18c66bc41807ff119b2",
@ -567,6 +583,14 @@
"markers": "python_version >= '3.8'", "markers": "python_version >= '3.8'",
"version": "==1.5.0" "version": "==1.5.0"
}, },
"pygments": {
"hashes": [
"sha256:61c16d2a8576dc0649d9f39e089b5f02bcd27fba10d8fb4dcc28173f7a45151f",
"sha256:9ea1544ad55cecf4b8242fab6dd35a93bbce657034b0611ee383099054ab6d8c"
],
"markers": "python_version >= '3.8'",
"version": "==2.19.1"
},
"pyparsing": { "pyparsing": {
"hashes": [ "hashes": [
"sha256:93d9577b88da0bbea8cc8334ee8b918ed014968fd2ec383e868fb8afb1ccef84", "sha256:93d9577b88da0bbea8cc8334ee8b918ed014968fd2ec383e868fb8afb1ccef84",
@ -592,6 +616,14 @@
"markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'", "markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'",
"version": "==2.9.0.post0" "version": "==2.9.0.post0"
}, },
"rich": {
"hashes": [
"sha256:439594978a49a09530cff7ebc4b5c7103ef57baf48d5ea3184f21d9a2befa098",
"sha256:6049d5e6ec054bf2779ab3358186963bac2ea89175919d699e378b99738c2a90"
],
"markers": "python_full_version >= '3.8.0'",
"version": "==13.9.4"
},
"scipy": { "scipy": {
"hashes": [ "hashes": [
"sha256:0c2f95de3b04e26f5f3ad5bb05e74ba7f68b837133a4492414b3afd79dfe540e", "sha256:0c2f95de3b04e26f5f3ad5bb05e74ba7f68b837133a4492414b3afd79dfe540e",
@ -641,6 +673,14 @@
"markers": "python_version >= '3.10'", "markers": "python_version >= '3.10'",
"version": "==1.14.1.6" "version": "==1.14.1.6"
}, },
"shellingham": {
"hashes": [
"sha256:7ecfff8f2fd72616f7481040475a65b2bf8af90a56c89140852d1120324e8686",
"sha256:8dbca0739d487e5bd35ab3ca4b36e11c4078f3a234bfce294b0a0291363404de"
],
"markers": "python_version >= '3.7'",
"version": "==1.5.4"
},
"simple-websocket": { "simple-websocket": {
"hashes": [ "hashes": [
"sha256:4af6069630a38ed6c561010f0e11a5bc0d4ca569b36306eb257cd9a192497c8c", "sha256:4af6069630a38ed6c561010f0e11a5bc0d4ca569b36306eb257cd9a192497c8c",
@ -657,6 +697,23 @@
"markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'", "markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'",
"version": "==1.17.0" "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": { "werkzeug": {
"hashes": [ "hashes": [
"sha256:54b78bf3716d19a65be4fceccc0d1d7b89e608834989dfae50ea87564639213e", "sha256:54b78bf3716d19a65be4fceccc0d1d7b89e608834989dfae50ea87564639213e",

View File

@ -1,11 +1,6 @@
import logging import logging
import multiprocessing as mp
from typing import NamedTuple
from . import ingest from . import cli
from . import config
from . import visualise
from . import aoa
logging.basicConfig( logging.basicConfig(
@ -13,35 +8,5 @@ logging.basicConfig(
format="%(asctime)s %(name)-40s %(levelname)-8s %(message)s", format="%(asctime)s %(name)-40s %(levelname)-8s %(message)s",
) )
if __name__ == "__main__":
class Receiver(NamedTuple): cli.app()
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()

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 socket
import multiprocessing as mp import multiprocessing as mp
import struct import struct
from typing import Callable
from datetime import datetime from datetime import datetime
import numpy as np import numpy as np
import numpy.typing as npt
from typing import NamedTuple
from .csi import CSI from .csi_frame import CSI
from .aoa import AoA from .. import config
from .preprocess import Preprocessor
from . import config
Host = tuple[str, int] Host = tuple[str, int]
CSICallback = Callable[[npt.NDArray[np.complex128]], None]
class FeitTransmitter: class FeitTransmitter:
@ -28,6 +30,8 @@ class FeitTransmitter:
) )
self.inject_server.send(b"stop\n") self.inject_server.send(b"stop\n")
self.inject_server.send(inject_start_string.encode()) 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: class FeitReceiver:
@ -46,9 +50,11 @@ class FeitReceiver:
self.logger = logging.getLogger( self.logger = logging.getLogger(
f"{__name__}.{self.__class__.__name__}-{self.host}" f"{__name__}.{self.__class__.__name__}-{self.host}"
) )
self.logger.info(f"Connected to {self.host}")
def listen(self, queue: "mp.Queue[CSI]"): def listen(self, queue: "mp.Queue[CSI]"):
prev_time = datetime.now() prev_time = datetime.now()
self.logger.info("Listening for CSI data")
while True: while True:
# This is the max size of a UDP packet. The size of the actual CSI # 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 # packet will depend on the frame format and channel width, which
@ -69,18 +75,12 @@ class CSIProcessor:
def __init__( def __init__(
self, self,
receiver_connections: dict[Host, "mp.Queue[CSI]"], receiver_connections: dict[Host, "mp.Queue[CSI]"],
webserver: "mp.Queue[AoA]",
): ):
self.pending_data: dict[Host, tuple[datetime, CSI]] = {} self.pending_data: dict[Host, tuple[datetime, CSI]] = {}
self.pending_data_lock = mp.Lock() self.pending_data_lock = mp.Lock()
self.logger = logging.getLogger(f"{__name__}.{self.__class__.__name__}") self.logger = logging.getLogger(f"{__name__}.{self.__class__.__name__}")
self.last_processed = datetime.now() self.last_processed = datetime.now()
self.preprocess = Preprocessor()
self.aoa = AoA()
self.connections = receiver_connections self.connections = receiver_connections
self.webserver = webserver
def add_data(self, host: Host, data: CSI): def add_data(self, host: Host, data: CSI):
if ( if (
@ -111,7 +111,7 @@ class CSIProcessor:
return False return False
return True return True
def process_data(self): def process_data(self, callback: CSICallback):
self.last_processed = datetime.now() self.last_processed = datetime.now()
antenna_data = [ antenna_data = [
np.expand_dims(self.pending_data[ip][1].matrix[:, antenna], axis=2) 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) all_data = np.concat(antenna_data, axis=1)
self.logger.info(f"Got final CSI data with shape {all_data.shape}") self.logger.info(f"Got final CSI data with shape {all_data.shape}")
processed = self.preprocess.preprocess(all_data) callback(all_data)
visualise.add_data(all_data, processed)
self.aoa.update(processed)
if not self.webserver.full():
self.webserver.put(self.aoa)
def process_forever(self): def process_forever(
self,
callback: CSICallback,
):
while True: while True:
for ip, queue in self.connections.items(): for ip, queue in self.connections.items():
while not queue.empty(): while not queue.empty():
self.add_data(ip, queue.get()) self.add_data(ip, queue.get())
if self.is_ready(): if self.is_ready():
self.process_data() self.process_data(callback)
else: else:
self.logger.debug("Not all data is ready") 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 as np
import numpy.typing as npt import numpy.typing as npt
from . import config from .. import config
from datetime import datetime from datetime import datetime
import logging import logging

View File

@ -1,7 +1,7 @@
import numpy as np import numpy as np
import logging import logging
from queue import Queue from queue import Queue
from . import config from .. import config
import numpy.typing as npt import numpy.typing as npt
from scipy.signal import correlate from scipy.signal import correlate
from scipy.signal import butter, sosfilt_zi, sosfilt from scipy.signal import butter, sosfilt_zi, sosfilt

View File

@ -11,7 +11,7 @@ import matplotlib.pyplot as plt
import io import io
import logging import logging
from ..aoa import AoA from ..processing.aoa import AoA
from .. import config from .. import config
import matplotlib import matplotlib