145 lines
4.9 KiB
Python
145 lines
4.9 KiB
Python
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)
|