import logging import time import socket import multiprocessing as mp import struct import torch from datetime import datetime import numpy as np from .csi import CSI from .aoa import AoA from .preprocess import Preprocessor from . import config device = torch.device("cuda" if torch.cuda.is_available() else "cpu") Host = tuple[str, int] class FeitTransmitter: def __init__(self): self.inject_server = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) self.inject_server.connect(config.INJECT_HOST) inject_start_string = ( f"feitcsi --frequency {config.CENTRAL_FREQUENCY_MHZ} " f"--channel-width {config.CHANNEL_WIDTH} " f"--format {config.FRAME_FORMAT} " f"--mode inject -s 1 --verbose " f"--inject-delay {1_000_000 // config.SAMPLE_RATE}" ) self.inject_server.send(b"stop\n") self.inject_server.send(inject_start_string.encode()) class FeitReceiver: def __init__(self, host: Host): self.host = host self.server = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) self.server.connect(host) self.start_string = ( f"feitcsi --frequency {config.CENTRAL_FREQUENCY_MHZ} " f"--channel-width {config.CHANNEL_WIDTH} " f"--format {config.FRAME_FORMAT} " f"--mode measure" ) self.server.send(b"stop\n") self.server.send(self.start_string.encode()) self.logger = logging.getLogger( f"{__name__}.{self.__class__.__name__}-{self.host}" ) 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 # packet will depend on the frame format and channel width, which # changes the number of subcarriers data = self.server.recv(65535) try: csidata = CSI(data) self.logger.debug( f"Received CSI data after {datetime.now() - prev_time}" ) prev_time = datetime.now() queue.put(csidata) except struct.error: self.logger.error("Failed to parse CSI data") 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 ( host in self.pending_data and self.last_processed < self.pending_data[host][0] ): self.logger.warning( f"Skipping data from {host} at {self.pending_data[host][0]}" ) 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 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): for host in config.RECEIVE_HOSTS: if ( host not in self.pending_data or self.pending_data[host][0] <= self.last_processed ): return False return True def process_data(self): self.last_processed = datetime.now() antenna_data = [ np.expand_dims(self.pending_data[ip][1].matrix[:, antenna], axis=2) for ip, antenna in config.ANTENNA_ORDER ] # We have data from all servers 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) processed_tensor = torch.tensor(processed, device=device) # visualise.add_data(all_data_tensor, processed) self.aoa.update(processed_tensor) self.logger.info("Processed data") if not self.webserver.full(): self.webserver.put(self.aoa) def process_forever(self): 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() else: self.logger.debug("Not all data is ready") time.sleep(0.001)