From 8e1a849b7c7e4c86b3c8a8710d3b3bf7415636b1 Mon Sep 17 00:00:00 2001 From: Christos Falas Date: Thu, 23 Jan 2025 11:13:17 +0000 Subject: [PATCH] format: re-format with ruff and isort --- pyproject.toml | 6 +++++ src/cli.py | 17 ++++++------- src/collection/csi_frame.py | 8 +++--- src/collection/ingest.py | 49 ++++++++++++++++++------------------ src/processing/aoa.py | 30 +++++++++++----------- src/processing/preprocess.py | 12 ++++----- 6 files changed, 65 insertions(+), 57 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index da9acff..884ab8c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,12 @@ [tool.ruff] line-length = 88 +[tool.ruff.lint] +extend-select = ["I", "E", "W", "F", "B", "Q", "ANN"] + +[tool.ruff.lint.pycodestyle] +max-doc-length = 88 + [tool.pyright] typeCheckingMode = "strict" reportMissingTypeStubs = "warning" diff --git a/src/cli.py b/src/cli.py index 40ae56e..40457cb 100644 --- a/src/cli.py +++ b/src/cli.py @@ -1,20 +1,19 @@ -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 +import typer + +from . import config, visualise +from .collection import ingest +from .processing.aoa import AoA +from .processing.preprocess import Preprocessor app = typer.Typer() @app.command() -def heatmap(): +def heatmap() -> None: preprocessor = Preprocessor() aoa = AoA() webapp_queue: "mp.Queue[AoA]" = mp.Queue(config.SAMPLE_RATE) @@ -22,7 +21,7 @@ def heatmap(): webapp = mp.Process(target=visualise.start, args=(webapp_queue,)) webapp.start() - def callback(antenna_data: npt.NDArray[np.complex128]): + def callback(antenna_data: npt.NDArray[np.complex128]) -> None: processed = preprocessor.preprocess(antenna_data) # visualise.add_data(all_data, processed) aoa.update(processed) diff --git a/src/collection/csi_frame.py b/src/collection/csi_frame.py index b984036..83291a9 100644 --- a/src/collection/csi_frame.py +++ b/src/collection/csi_frame.py @@ -1,8 +1,8 @@ import struct + import numpy as np import numpy.typing as npt - RATE_MCS_MOD_TYPE_POS = 8 RATE_MCS_MOD_TYPE_MSK = 0x7 << RATE_MCS_MOD_TYPE_POS RATE_MCS_CCK_MSK = 0 << RATE_MCS_MOD_TYPE_POS @@ -37,7 +37,7 @@ RATE_MCS_BEAMF_MSK = 1 << RATE_MCS_BEAMF_POS class CSIHeader: - def __init__(self, data: bytes): + def __init__(self, data: bytes) -> None: self.csi_size = struct.unpack("I", data[0:4])[0] self.ftm_clock = struct.unpack("I", data[8:12])[0] self.num_rx = data[46] @@ -93,7 +93,7 @@ class CSIHeader: class CSI: @staticmethod - def parseCsiData(data: bytes, header: CSIHeader): + def parseCsiData(data: bytes, header: CSIHeader) -> npt.NDArray[np.complex128]: csi_matrix: npt.NDArray[np.complex128] = np.zeros( ( header.num_subcarriers, @@ -113,7 +113,7 @@ class CSI: return csi_matrix - def __init__(self, data: bytes): + def __init__(self, data: bytes) -> None: self.header = CSIHeader(data[:272]) self.matrix = self.parseCsiData( data[272 : 272 + self.header.csi_size], self.header diff --git a/src/collection/ingest.py b/src/collection/ingest.py index 31c6030..016ed0f 100644 --- a/src/collection/ingest.py +++ b/src/collection/ingest.py @@ -1,26 +1,25 @@ import logging -import time -import socket import multiprocessing as mp -import threading -import subprocess +import socket import struct -from typing import Callable +import subprocess +import threading +import time from datetime import datetime +from typing import Callable, NamedTuple, NoReturn + import numpy as np import numpy.typing as npt -from typing import NamedTuple -from .csi_frame import CSI from .. import config - +from .csi_frame import CSI Host = tuple[str, int] CSICallback = Callable[[npt.NDArray[np.complex128]], None] class FeitHost: - def __init__(self, host: Host, command: str): + def __init__(self, host: Host, command: str) -> None: self.command = command self.host = host self.logger = logging.getLogger( @@ -29,7 +28,7 @@ class FeitHost: self.checker = threading.Thread(target=self.check_continuous) self.checker.start() - def check_connection(self): + def check_connection(self) -> bool: feitcsi_status = subprocess.run( f"ssh root@{self.host[0]} pgrep feitcsi", check=False, @@ -38,18 +37,20 @@ class FeitHost: ) return feitcsi_status.returncode == 0 - def connect(self): + def connect(self) -> None: self.server = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) self.server.connect(self.host) self.server.send(b"stop\n") self.server.send(self.command.encode()) self.logger.info(f"Connected to {self.host}") - def check_continuous(self): + def check_continuous(self) -> NoReturn: """ Repeatedly check if the FeitCSI service is running - Because FeitCSI is using TCP, we will get no information if the service stops, or the computer is not reachable. In order to make debugging easier, this checks and logs continuously if the service is running. + Because FeitCSI is using TCP, we will get no information if the service stops, + or the computer is not reachable. In order to make debugging easier, this checks + and logs continuously if the service is running. """ last_status = False while True: @@ -64,7 +65,7 @@ class FeitHost: class FeitTransmitter(FeitHost): - def __init__(self): + def __init__(self) -> None: command = ( f"feitcsi --frequency {config.CENTRAL_FREQUENCY_MHZ} " f"--channel-width {config.CHANNEL_WIDTH} " @@ -76,7 +77,7 @@ class FeitTransmitter(FeitHost): class FeitReceiver(FeitHost): - def __init__(self, host: Host): + def __init__(self, host: Host) -> None: command = ( f"feitcsi --frequency {config.CENTRAL_FREQUENCY_MHZ} " f"--channel-width {config.CHANNEL_WIDTH} " @@ -85,7 +86,7 @@ class FeitReceiver(FeitHost): ) super().__init__(host, command) - def listen(self, queue: "mp.Queue[CSI]"): + def listen(self, queue: "mp.Queue[CSI]") -> NoReturn: prev_time = datetime.now() self.logger.info("Listening for CSI data") while True: @@ -110,14 +111,14 @@ class CSIProcessor: def __init__( self, receiver_connections: dict[Host, "mp.Queue[CSI]"], - ): + ) -> None: 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.connections = receiver_connections - def add_data(self, host: Host, data: CSI): + def add_data(self, host: Host, data: CSI) -> None: if ( host in self.pending_data and self.last_processed < self.pending_data[host][0] @@ -129,15 +130,15 @@ class CSIProcessor: with self.pending_data_lock: self.pending_data[host] = (datetime.now(), data) - # Useful for figuring out the correct antenna order - RSSI values will decrease - # when the specific antenna is disconnected + # Useful for figuring out the correct antenna order - RSSI values will + # decrease when the specific antenna is disconnected rssis = [ (ip, csi.header.rssi1, csi.header.rssi2) for ip, (_, csi) in sorted(self.pending_data.items()) ] self.logger.debug("Antenna RSSI values: {}".format(rssis)) - def is_ready(self): + def is_ready(self) -> bool: for host in config.RECEIVE_HOSTS: if ( host not in self.pending_data @@ -146,7 +147,7 @@ class CSIProcessor: return False return True - def process_data(self, callback: CSICallback): + def process_data(self, callback: CSICallback) -> None: self.last_processed = datetime.now() antenna_data = [ np.expand_dims(self.pending_data[ip][1].matrix[:, antenna], axis=2) @@ -162,7 +163,7 @@ class CSIProcessor: def process_forever( self, callback: CSICallback, - ): + ) -> NoReturn: while True: for ip, queue in self.connections.items(): while not queue.empty(): @@ -180,7 +181,7 @@ class Receiver(NamedTuple): queue: "mp.Queue[CSI]" -def start_processing(csi_callback: CSICallback): +def start_processing(csi_callback: CSICallback) -> None: receivers = [ Receiver(ip, FeitReceiver(ip), mp.Queue(config.SAMPLE_RATE)) for ip in config.RECEIVE_HOSTS diff --git a/src/processing/aoa.py b/src/processing/aoa.py index 9f427d9..139185a 100644 --- a/src/processing/aoa.py +++ b/src/processing/aoa.py @@ -1,22 +1,23 @@ -import numpy as np -import numpy.typing as npt -from .. import config +import logging from datetime import datetime -import logging +import numpy as np +import numpy.typing as npt + +from .. import config logger = logging.getLogger(__name__) class AoA: - def __init__(self): + def __init__(self) -> None: self.historical_autocorr = np.array([]) self.N_subcarriers = -1 self.N_rx = -1 self.timestamp = datetime.now() pass - def smooth(self, data: npt.NDArray[np.complex128]): + def smooth(self, data: npt.NDArray[np.complex128]) -> npt.NDArray[np.complex128]: assert len(data.shape) == 3 M = data.shape[0] # Number of subcarriers @@ -44,7 +45,7 @@ class AoA: return H_sm - def update(self, data: npt.NDArray[np.complex128]): + def update(self, data: npt.NDArray[np.complex128]) -> None: self.timestamp = datetime.now() H_sm = self.smooth(data) @@ -76,9 +77,11 @@ class AoA: 2j * np.pi * config.CENTRAL_FREQUENCY_HZ * config.ANTENNA_SPACING / config.C ) - def steering_vector(self, theta: float, tof: float): - omega_t = np.exp(-2j * np.pi * config.DELTA_F * tof) - phi_theta = np.exp( + def steering_vector( + self, theta: float, tof: float + ) -> npt.NDArray[np.complexfloating]: + omega_t: npt.NDArray[np.complex128] = np.exp(-2j * np.pi * config.DELTA_F * tof) + phi_theta: npt.NDArray[np.complex128] = np.exp( 2j * np.pi * config.CENTRAL_FREQUENCY_HZ @@ -96,7 +99,7 @@ class AoA: steering = antenna_v * phis return steering.T.reshape(-1) - def evaluate(self, theta: float, tof: float): + def evaluate(self, theta: float, tof: float) -> float: try: steering = self.steering_vector(theta, tof) steering_h = np.conj(steering).T @@ -109,7 +112,7 @@ class AoA: return np.abs(c.real) -def test_smoothing(): +def test_smoothing() -> None: row, col = np.indices((4, 2)) data = row + 1j * col aoa = AoA() @@ -120,10 +123,9 @@ def test_smoothing(): expected = np.hstack([H_01, H_12]) print(expected) assert np.allclose(smoothed, expected) - pass -def test_steering_vector(): +def test_steering_vector() -> None: aoa = AoA() aoa.N_subcarriers = 10 aoa.N_rx = 2 diff --git a/src/processing/preprocess.py b/src/processing/preprocess.py index f3a290c..96b925c 100644 --- a/src/processing/preprocess.py +++ b/src/processing/preprocess.py @@ -1,11 +1,11 @@ -import numpy as np import logging from queue import Queue -from .. import config -import numpy.typing as npt -from scipy.signal import correlate -from scipy.signal import butter, sosfilt_zi, sosfilt +import numpy as np +import numpy.typing as npt +from scipy.signal import butter, correlate, sosfilt, sosfilt_zi + +from .. import config logger = logging.getLogger(__name__) logger.setLevel(logging.DEBUG) @@ -14,7 +14,7 @@ np.seterr(invalid="ignore") class Preprocessor: - def __init__(self): + def __init__(self) -> None: self.prev_entries: Queue[npt.NDArray[np.complex128]] = Queue(maxsize=100) self.short_term_avg = np.zeros((1,), dtype=np.complex128) self.long_term_avg = np.zeros((1,), dtype=np.complex128)