dissertation/src/ingest.py
2025-01-24 12:04:05 +00:00

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)