Files
SDR-Rover/tests/test_bpsk_radio.py
2026-08-19 18:14:13 +03:00

369 lines
14 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Проверки protocol/bpsk_radio.py.
Модуль был собран переносом функций из Lab018, Lab019 и Lab023. Главная
задача этих проверок — доказать, что перенос ничего не изменил: результаты
сравниваются с оригиналами, загруженными из лабораторных в изоляции, без
выполнения их модульного кода.
"""
from __future__ import annotations
import ast
import math
import pathlib
import struct
import numpy as np
import pytest
from protocol import bpsk_radio as radio
from protocol.packet import MESSAGE_TYPE_IMAGE_FRAGMENT, build_packet, parse_packet
PROJECT_ROOT = pathlib.Path(__file__).resolve().parent.parent
SHARED_CONSTANTS = {
"RADIO_SYNC_WORD",
"PREAMBLE_BIT_COUNT",
"SYMBOL_RATE",
"CFO_REFINEMENT_HALF_WIDTH_HZ",
"CFO_REFINEMENT_STEP_HZ",
}
def _load_isolated(relative_path: str, function_names: set[str]) -> dict:
"""Выполнить только заданные функции лабораторной, без кода модуля.
Лабораторные при импорте создают каталоги и пишут файлы, поэтому
импортировать их в тестах нельзя.
"""
path = PROJECT_ROOT / relative_path
tree = ast.parse(path.read_text(encoding="utf-8"))
kept = [
node
for node in tree.body
if (isinstance(node, ast.FunctionDef) and node.name in function_names)
or (isinstance(node, ast.Assign) and ast.unparse(node.targets[0]) in SHARED_CONSTANTS)
]
namespace: dict = {"np": np, "struct": struct}
exec(compile(ast.Module(body=kept, type_ignores=[]), str(path), "exec"), namespace)
return namespace
@pytest.fixture(scope="module")
def original_lab023() -> dict:
return _load_isolated(
"experiments/lab023_guarded_cfo_correction.py",
{
"bytes_to_bits",
"bits_to_bytes",
"bpsk_modulate",
"bpsk_demodulate",
"estimate_carrier_parameters",
"correct_phase_and_frequency",
},
)
@pytest.fixture(scope="module")
def original_lab019() -> dict:
return _load_isolated(
"experiments/lab019_bpsk_receiver.py",
{
"build_frame_marker",
"find_radio_frame",
"bytes_to_bits",
"bits_to_bytes",
"bpsk_modulate",
"bpsk_demodulate",
"root_raised_cosine_taps",
},
)
@pytest.fixture(scope="module")
def original_lab018() -> dict:
return _load_isolated("experiments/lab018_bpsk_radio_frame.py", {"root_raised_cosine_taps"})
def _identical(left, right) -> bool:
if isinstance(left, np.ndarray) or isinstance(right, np.ndarray):
return np.array_equal(np.asarray(left), np.asarray(right))
if isinstance(left, dict):
return sorted(left) == sorted(right) and all(_identical(left[k], right[k]) for k in left)
if isinstance(left, (tuple, list)):
return len(left) == len(right) and all(_identical(a, b) for a, b in zip(left, right))
return left == right
# ---------------------------------------------------------------- эквивалентность
@pytest.mark.parametrize("length", [1, 7, 64, 255])
def test_bit_conversion_matches_lab023(original_lab023: dict, length: int) -> None:
data = bytes(np.random.default_rng(length).integers(0, 256, size=length, dtype=np.uint8))
ours = radio.bytes_to_bits(data)
assert _identical(ours, original_lab023["bytes_to_bits"](data))
assert _identical(radio.bits_to_bytes(ours), original_lab023["bits_to_bytes"](ours))
@pytest.mark.parametrize("length", [8, 64, 256])
def test_modulation_matches_lab023(original_lab023: dict, length: int) -> None:
bits = np.random.default_rng(length).integers(0, 2, size=length).astype(np.uint8)
ours = radio.bpsk_modulate(bits)
assert _identical(ours, original_lab023["bpsk_modulate"](bits))
assert _identical(radio.bpsk_demodulate(ours), original_lab023["bpsk_demodulate"](ours))
@pytest.mark.parametrize("rolloff,sps,span", [(0.35, 32, 10), (0.35, 128, 10), (0.5, 64, 8)])
def test_shaping_filter_matches_lab018(original_lab018: dict, rolloff: float, sps: int, span: int) -> None:
assert _identical(
radio.root_raised_cosine_taps(rolloff, sps, span),
original_lab018["root_raised_cosine_taps"](rolloff, sps, span),
)
def test_frame_marker_matches_lab019(original_lab019: dict) -> None:
"""Маркер возвращает пару: биты и символы. Совпасть должны обе."""
assert _identical(radio.build_frame_marker(), original_lab019["build_frame_marker"]())
@pytest.mark.parametrize("seed", [1, 2, 3])
def test_carrier_estimation_matches_lab023(original_lab023: dict, seed: int) -> None:
"""Оценка ухода частоты обязана совпасть при произвольных фазе и уходе."""
rng = np.random.default_rng(seed)
_, marker = radio.build_frame_marker()
phase = float(rng.uniform(-3.0, 3.0))
increment = float(rng.uniform(-0.005, 0.005))
impaired = marker * np.exp(1j * (phase + increment * np.arange(len(marker))))
impaired = impaired + 0.02 * rng.standard_normal(len(marker))
assert _identical(
radio.estimate_carrier_parameters(impaired, marker),
original_lab023["estimate_carrier_parameters"](impaired, marker),
)
assert _identical(
radio.correct_phase_and_frequency(impaired, phase, increment),
original_lab023["correct_phase_and_frequency"](impaired, phase, increment),
)
# ---------------------------------------------------------------- новые функции кадра
@pytest.mark.parametrize("payload_size", [1, 17, 512, radio.MAXIMUM_PROTOCOL_PACKET_BYTES])
def test_frame_roundtrip(payload_size: int) -> None:
packet = bytes(np.random.default_rng(payload_size).integers(0, 256, size=payload_size, dtype=np.uint8))
bits, symbols, marker = radio.build_radio_frame(packet)
assert len(bits) == radio.radio_frame_bit_count(payload_size)
assert len(symbols) == len(bits)
assert len(marker) == radio.MARKER_BIT_COUNT
assert radio.parse_radio_frame(bits) == packet
def test_frame_rejects_sizes_outside_range() -> None:
with pytest.raises(ValueError):
radio.build_radio_frame(b"")
with pytest.raises(ValueError):
radio.build_radio_frame(b"\x00" * (radio.MAXIMUM_PROTOCOL_PACKET_BYTES + 1))
with pytest.raises(TypeError):
radio.build_radio_frame("не байты")
def test_frame_with_broken_sync_word_is_rejected() -> None:
bits, _, _ = radio.build_radio_frame(b"payload")
broken = bits.copy()
broken[radio.PREAMBLE_BIT_COUNT] ^= 1
assert radio.parse_radio_frame(broken) is None
def test_truncated_frame_is_rejected() -> None:
bits, _, _ = radio.build_radio_frame(b"payload")
assert radio.parse_radio_frame(bits[:10]) is None
assert radio.parse_radio_frame(bits[:-8]) is None
def test_frame_carries_a_protocol_packet() -> None:
"""Кадр должен переносить настоящий пакет так, чтобы CRC32 сошлась."""
packet = build_packet(b"fragment payload", MESSAGE_TYPE_IMAGE_FRAGMENT, 5)
bits, _, _ = radio.build_radio_frame(packet)
assert parse_packet(radio.parse_radio_frame(bits)).payload == b"fragment payload"
def test_frame_search_finds_the_start_in_a_shaped_signal() -> None:
"""Сквозная проверка: формирование, фильтрация и поиск кадра."""
samples_per_symbol = 32
taps = radio.root_raised_cosine_taps(0.35, samples_per_symbol, 10)
packet = build_packet(b"HELLO SDR", MESSAGE_TYPE_IMAGE_FRAGMENT, 1)
bits, symbols, marker = radio.build_radio_frame(packet)
upsampled = np.zeros(len(symbols) * samples_per_symbol, dtype=complex)
upsampled[::samples_per_symbol] = symbols
matched = np.convolve(np.convolve(upsampled, taps, mode="full"), taps, mode="full")
found = radio.find_radio_frame(matched, marker, samples_per_symbol)
start = found["start_symbol_index"]
recovered = found["symbol_samples"][start : start + len(bits)]
assert radio.parse_radio_frame(radio.bpsk_demodulate(recovered)) == packet
# ------------------------------------------------------ Lab043: общие примитивы
def test_known_tone_peak_search_finds_both_windows() -> None:
frequencies = np.linspace(-100_000.0, 100_000.0, 4001)
powers = np.ones_like(frequencies)
powers[np.argmin(np.abs(frequencies + 51_200.0))] = 100.0
powers[np.argmin(np.abs(frequencies - 49_300.0))] = 80.0
peaks = radio.find_known_tone_peaks(frequencies, powers, 50_000.0, 30_000.0)
assert math.isclose(peaks["low_frequency_hz"], -51_200.0, abs_tol=1.0)
assert math.isclose(peaks["high_frequency_hz"], 49_300.0, abs_tol=1.0)
assert peaks["low_power"] == 100.0
assert peaks["high_power"] == 80.0
def test_known_tone_peak_search_returns_nan_without_bins() -> None:
frequencies = np.linspace(-1_000.0, 1_000.0, 101)
powers = np.ones_like(frequencies)
peaks = radio.find_known_tone_peaks(frequencies, powers, 50_000.0, 1_000.0)
assert math.isnan(peaks["low_frequency_hz"])
assert math.isnan(peaks["high_frequency_hz"])
@pytest.mark.parametrize("ppm", [20.0, -20.0, 100.0, -100.0])
def test_two_tone_clock_estimate_has_correct_sign_and_magnitude(ppm: float) -> None:
scale = 1.0 + ppm * 1e-6
carrier_offset_hz = -1_250.0
low = carrier_offset_hz - 50_000.0 * scale
high = carrier_offset_hz + 50_000.0 * scale
estimate = radio.estimate_two_tone_offsets(low, high, 50_000.0)
assert math.isclose(estimate["carrier_offset_hz"], carrier_offset_hz, abs_tol=1e-9)
assert math.isclose(estimate["clock_scale"], scale, abs_tol=1e-12)
assert math.isclose(estimate["sample_clock_error_ppm"], ppm, abs_tol=1e-6)
def test_two_tone_invalid_estimate_is_nan_not_zero() -> None:
estimate = radio.estimate_two_tone_offsets(float("nan"), 50_000.0, 50_000.0)
assert all(math.isnan(value) for value in estimate.values())
@pytest.mark.parametrize("carrier_offset_hz", [730.0, -730.0])
def test_coarse_frequency_correction_handles_both_signs(carrier_offset_hz: float) -> None:
sample_rate_hz = 20_000.0
indexes = np.arange(20_000, dtype=np.float64)
impaired = np.exp(1j * 2.0 * np.pi * carrier_offset_hz * indexes / sample_rate_hz)
corrected = radio.apply_coarse_frequency_correction(
impaired,
carrier_offset_hz,
sample_rate_hz,
)
assert np.max(np.abs(corrected - 1.0)) < 1e-9
@pytest.mark.parametrize(
"carrier_offset_hz",
[-10.0, -5.0, -2.0, -1.2, -1.0, -0.5, 0.5, 1.0, 1.2, 2.0, 5.0, 10.0],
)
def test_known_pilot_estimator_resolves_sub_hertz_cfo_with_hardware_like_phase_noise(
carrier_offset_hz: float,
) -> None:
symbol_count = 1_280
indexes = np.arange(symbol_count, dtype=np.float64)
known = np.where((indexes.astype(np.int64) * 73 + 19) % 127 < 64, 1.0, -1.0).astype(
np.complex128
)
estimates: list[float] = []
valid_flags: list[bool] = []
for repetition in range(96):
random_generator = np.random.default_rng(
43_000_000
+ int(round((carrier_offset_hz + 20.0) * 1_000.0))
+ repetition
)
phase_noise = random_generator.normal(0.0, 0.4691, symbol_count)
received = known * np.exp(
1j
* (
2.0 * np.pi * carrier_offset_hz * indexes / radio.SYMBOL_RATE
+ phase_noise
)
)
estimate = radio.estimate_known_pilot_carrier(received, known)
valid_flags.append(estimate.valid)
estimates.append(estimate.frequency_hz)
errors = np.asarray(estimates) - carrier_offset_hz
assert all(valid_flags)
assert np.all(np.sign(estimates) == np.sign(carrier_offset_hz))
assert abs(float(np.mean(errors))) < 0.08
assert float(np.std(errors, ddof=1)) < 0.18
assert float(np.percentile(np.abs(errors), 95.0)) < 0.35
def test_known_pilot_estimator_rejects_noise_instead_of_reporting_false_cfo() -> None:
random_generator = np.random.default_rng(43_043)
known = np.resize(np.asarray([-1.0, 1.0], dtype=np.complex128), 1_280)
noise = (
random_generator.normal(0.0, 1.0, len(known))
+ 1j * random_generator.normal(0.0, 1.0, len(known))
)
estimate = radio.estimate_known_pilot_carrier(noise, known)
assert not estimate.valid
assert estimate.invalid_reason
assert math.isnan(estimate.frequency_hz)
assert math.isnan(estimate.phase_increment_rad_per_symbol)
def _sample_at_positions(signal: np.ndarray, positions: np.ndarray) -> np.ndarray:
indexes = np.arange(len(signal), dtype=np.float64)
real = np.interp(positions, indexes, signal.real)
imaginary = np.interp(positions, indexes, signal.imag)
return real + 1j * imaginary
@pytest.mark.parametrize("ppm", [20.0, -20.0, 100.0, -100.0])
def test_clock_resampling_direction_reduces_timing_error(ppm: float) -> None:
scale = 1.0 + ppm * 1e-6
sample_count = 200_000
indexes = np.arange(sample_count, dtype=np.float64)
reference = np.exp(1j * 2.0 * np.pi * 0.071 * indexes)
received_length = math.floor(sample_count / scale)
received_positions = np.arange(received_length, dtype=np.float64) * scale
received = _sample_at_positions(reference, received_positions)
corrected = radio.resample_for_clock_scale(received, scale)
wrong_direction = radio.resample_for_clock_scale(received, 1.0 / scale)
uncorrected_count = min(len(received), len(reference))
corrected_count = min(len(corrected), len(reference)) - 2
wrong_count = min(len(wrong_direction), len(reference)) - 2
uncorrected_error = float(
np.mean(np.abs(received[:uncorrected_count] - reference[:uncorrected_count]) ** 2)
)
corrected_error = float(
np.mean(np.abs(corrected[:corrected_count] - reference[:corrected_count]) ** 2)
)
wrong_error = float(
np.mean(np.abs(wrong_direction[:wrong_count] - reference[:wrong_count]) ** 2)
)
assert len(corrected) == round(len(received) * scale)
assert corrected_error < uncorrected_error
assert corrected_error < wrong_error