614 lines
19 KiB
Python
614 lines
19 KiB
Python
"""
|
||
Systematic packet-erasure FEC over GF(256) for Lab028 wire packets.
|
||
|
||
The code uses a Vandermonde generator matrix over GF(256), transformed to
|
||
systematic form. The primitive polynomial is x^8+x^4+x^3+x^2+1 (0x11D).
|
||
Any k available symbols from an n=k+r block reconstruct all k source
|
||
symbols.
|
||
|
||
Each source or parity symbol is carried in a fixed 32-byte outer header::
|
||
|
||
!4sBBHIHHHHHHII
|
||
|
||
Offset Size Field
|
||
0 4 magic (b"SFE1")
|
||
4 1 version
|
||
5 1 flags (bit 0: parity)
|
||
6 2 header_size
|
||
8 4 block_id
|
||
12 2 symbol_index
|
||
14 2 source_count (k)
|
||
16 2 parity_count (r)
|
||
18 2 symbol_size
|
||
20 2 data_length
|
||
22 2 reserved16 (zero)
|
||
24 4 reserved32 (zero)
|
||
28 4 outer_crc32
|
||
|
||
Source symbols carry the original serialized Lab028 packet without padding.
|
||
Parity symbols carry symbol_size bytes. Zero padding is applied only during
|
||
GF(256) mathematics. A recovered source symbol is trimmed from the Lab028
|
||
payload_length field and must pass the original Lab028 packet CRC32.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from dataclasses import dataclass
|
||
from functools import lru_cache
|
||
import struct
|
||
import zlib
|
||
|
||
from protocol.video_packet import (
|
||
HEADER_FORMAT as INNER_HEADER_FORMAT,
|
||
HEADER_SIZE as INNER_HEADER_SIZE,
|
||
MAGIC as INNER_MAGIC,
|
||
VERSION as INNER_VERSION,
|
||
decode_packet as decode_inner_packet,
|
||
)
|
||
|
||
|
||
GF_PRIMITIVE_POLYNOMIAL = 0x11D
|
||
GF_FIELD_SIZE = 256
|
||
GF_ORDER = 255
|
||
|
||
OUTER_MAGIC = b"SFE1"
|
||
OUTER_VERSION = 1
|
||
OUTER_FLAG_PARITY = 0x01
|
||
OUTER_HEADER_FORMAT = "!4sBBHIHHHHHHII"
|
||
OUTER_HEADER_SIZE = struct.calcsize(OUTER_HEADER_FORMAT)
|
||
MAX_SYMBOL_SIZE = 0xFFFF
|
||
MAX_BLOCK_SYMBOLS = 0xFF
|
||
|
||
|
||
class FECError(ValueError):
|
||
"""Base class for outer packets and erasure-code failures."""
|
||
|
||
|
||
class OuterPacketCRCError(FECError):
|
||
"""The outer packet failed its CRC32 check."""
|
||
|
||
|
||
class InsufficientSymbolsError(FECError):
|
||
"""Fewer than k unique symbols are available."""
|
||
|
||
|
||
class MatrixSingularError(FECError):
|
||
"""A GF(256) matrix has no inverse."""
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class OuterSymbol:
|
||
"""Decoded source or parity symbol after outer CRC validation."""
|
||
|
||
block_id: int
|
||
symbol_index: int
|
||
source_count: int
|
||
parity_count: int
|
||
symbol_size: int
|
||
data: bytes
|
||
is_parity: bool
|
||
outer_crc32: int = 0
|
||
|
||
|
||
@dataclass(frozen=True)
|
||
class DecodedFECBlock:
|
||
"""All reconstructed Lab028 packets from one FEC block."""
|
||
|
||
block_id: int
|
||
source_packets: tuple[bytes, ...]
|
||
recovered_indices: tuple[int, ...]
|
||
|
||
|
||
def _build_gf_tables() -> tuple[
|
||
tuple[int, ...],
|
||
tuple[int, ...],
|
||
]:
|
||
exponent = [0] * (GF_ORDER * 2)
|
||
logarithm = [0] * GF_FIELD_SIZE
|
||
value = 1
|
||
for index in range(GF_ORDER):
|
||
exponent[index] = value
|
||
logarithm[value] = index
|
||
value <<= 1
|
||
if value & GF_FIELD_SIZE:
|
||
value ^= GF_PRIMITIVE_POLYNOMIAL
|
||
for index in range(GF_ORDER, GF_ORDER * 2):
|
||
exponent[index] = exponent[index - GF_ORDER]
|
||
return tuple(exponent), tuple(logarithm)
|
||
|
||
|
||
GF_EXP, GF_LOG = _build_gf_tables()
|
||
|
||
|
||
def gf_add(left: int, right: int) -> int:
|
||
"""Addition/subtraction in GF(256)."""
|
||
|
||
return left ^ right
|
||
|
||
|
||
def gf_mul(left: int, right: int) -> int:
|
||
"""Multiplication in GF(256)."""
|
||
|
||
if left == 0 or right == 0:
|
||
return 0
|
||
return GF_EXP[GF_LOG[left] + GF_LOG[right]]
|
||
|
||
|
||
def gf_div(dividend: int, divisor: int) -> int:
|
||
"""Division in GF(256)."""
|
||
|
||
if divisor == 0:
|
||
raise ZeroDivisionError("GF(256) division by zero")
|
||
if dividend == 0:
|
||
return 0
|
||
return GF_EXP[
|
||
(GF_LOG[dividend] - GF_LOG[divisor]) % GF_ORDER
|
||
]
|
||
|
||
|
||
def gf_inverse(value: int) -> int:
|
||
"""Multiplicative inverse in GF(256)."""
|
||
|
||
if value == 0:
|
||
raise ZeroDivisionError("zero has no GF(256) inverse")
|
||
return GF_EXP[GF_ORDER - GF_LOG[value]]
|
||
|
||
|
||
def gf_pow(value: int, exponent: int) -> int:
|
||
"""Non-negative integer power in GF(256)."""
|
||
|
||
if exponent < 0:
|
||
raise ValueError("negative GF exponent is unsupported")
|
||
if exponent == 0:
|
||
return 1
|
||
if value == 0:
|
||
return 0
|
||
return GF_EXP[(GF_LOG[value] * exponent) % GF_ORDER]
|
||
|
||
|
||
GF_MUL_TRANSLATIONS = tuple(
|
||
bytes(gf_mul(coefficient, value) for value in range(256))
|
||
for coefficient in range(256)
|
||
)
|
||
|
||
|
||
def matrix_multiply(
|
||
left: tuple[tuple[int, ...], ...],
|
||
right: tuple[tuple[int, ...], ...],
|
||
) -> tuple[tuple[int, ...], ...]:
|
||
"""Multiply two small matrices over GF(256)."""
|
||
|
||
if not left or not right:
|
||
raise ValueError("matrices must not be empty")
|
||
inner_size = len(left[0])
|
||
if inner_size != len(right):
|
||
raise ValueError("matrix dimensions do not match")
|
||
column_count = len(right[0])
|
||
result = []
|
||
for row in left:
|
||
if len(row) != inner_size:
|
||
raise ValueError("left matrix is ragged")
|
||
result_row = []
|
||
for column_index in range(column_count):
|
||
value = 0
|
||
for inner_index in range(inner_size):
|
||
value ^= gf_mul(
|
||
row[inner_index],
|
||
right[inner_index][column_index],
|
||
)
|
||
result_row.append(value)
|
||
result.append(tuple(result_row))
|
||
return tuple(result)
|
||
|
||
|
||
def matrix_inverse(
|
||
matrix: tuple[tuple[int, ...], ...],
|
||
) -> tuple[tuple[int, ...], ...]:
|
||
"""Invert a square matrix using GF(256) Gauss-Jordan elimination."""
|
||
|
||
size = len(matrix)
|
||
if size == 0 or any(len(row) != size for row in matrix):
|
||
raise ValueError("matrix must be non-empty and square")
|
||
augmented = [
|
||
list(row)
|
||
+ [1 if row_index == column_index else 0
|
||
for column_index in range(size)]
|
||
for row_index, row in enumerate(matrix)
|
||
]
|
||
|
||
for pivot_column in range(size):
|
||
pivot_row = next(
|
||
(
|
||
row_index
|
||
for row_index in range(pivot_column, size)
|
||
if augmented[row_index][pivot_column] != 0
|
||
),
|
||
None,
|
||
)
|
||
if pivot_row is None:
|
||
raise MatrixSingularError("GF(256) matrix is singular")
|
||
if pivot_row != pivot_column:
|
||
augmented[pivot_column], augmented[pivot_row] = (
|
||
augmented[pivot_row],
|
||
augmented[pivot_column],
|
||
)
|
||
|
||
pivot_inverse = gf_inverse(
|
||
augmented[pivot_column][pivot_column]
|
||
)
|
||
augmented[pivot_column] = [
|
||
gf_mul(value, pivot_inverse)
|
||
for value in augmented[pivot_column]
|
||
]
|
||
|
||
for row_index in range(size):
|
||
if row_index == pivot_column:
|
||
continue
|
||
factor = augmented[row_index][pivot_column]
|
||
if factor == 0:
|
||
continue
|
||
augmented[row_index] = [
|
||
value ^ gf_mul(factor, pivot_value)
|
||
for value, pivot_value in zip(
|
||
augmented[row_index],
|
||
augmented[pivot_column],
|
||
)
|
||
]
|
||
|
||
return tuple(
|
||
tuple(row[size:]) for row in augmented
|
||
)
|
||
|
||
|
||
@lru_cache(maxsize=None)
|
||
def systematic_generator_matrix(
|
||
source_count: int,
|
||
parity_count: int,
|
||
) -> tuple[tuple[int, ...], ...]:
|
||
"""Return an n×k systematic Vandermonde generator matrix."""
|
||
|
||
if not 1 <= source_count <= MAX_BLOCK_SYMBOLS:
|
||
raise ValueError("source_count is outside 1...255")
|
||
if not 0 <= parity_count <= (
|
||
MAX_BLOCK_SYMBOLS - source_count
|
||
):
|
||
raise ValueError("source_count + parity_count exceeds 255")
|
||
total_count = source_count + parity_count
|
||
vandermonde = tuple(
|
||
tuple(
|
||
gf_pow(row_index + 1, column_index)
|
||
for column_index in range(source_count)
|
||
)
|
||
for row_index in range(total_count)
|
||
)
|
||
top_inverse = matrix_inverse(
|
||
vandermonde[:source_count]
|
||
)
|
||
generator = matrix_multiply(vandermonde, top_inverse)
|
||
identity = tuple(
|
||
tuple(
|
||
1 if row_index == column_index else 0
|
||
for column_index in range(source_count)
|
||
)
|
||
for row_index in range(source_count)
|
||
)
|
||
if generator[:source_count] != identity:
|
||
raise RuntimeError("generator matrix is not systematic")
|
||
return generator
|
||
|
||
|
||
def _linear_combine(
|
||
coefficients: tuple[int, ...],
|
||
symbols: tuple[bytes, ...],
|
||
symbol_size: int,
|
||
) -> bytes:
|
||
"""Combine equal-size byte strings over GF(256), using C-level helpers."""
|
||
|
||
accumulator = 0
|
||
for coefficient, symbol in zip(coefficients, symbols):
|
||
if coefficient == 0:
|
||
continue
|
||
translated = symbol.translate(
|
||
GF_MUL_TRANSLATIONS[coefficient]
|
||
)
|
||
accumulator ^= int.from_bytes(translated, "little")
|
||
return accumulator.to_bytes(symbol_size, "little")
|
||
|
||
|
||
def encode_parity_symbols(
|
||
source_symbols: tuple[bytes, ...],
|
||
parity_count: int,
|
||
) -> tuple[bytes, ...]:
|
||
"""Encode zero-padded source symbols into parity symbols."""
|
||
|
||
if not source_symbols:
|
||
raise ValueError("source symbol list is empty")
|
||
symbol_size = len(source_symbols[0])
|
||
if symbol_size == 0:
|
||
raise ValueError("symbol size is zero")
|
||
if any(len(symbol) != symbol_size for symbol in source_symbols):
|
||
raise ValueError("source symbols must have equal size")
|
||
generator = systematic_generator_matrix(
|
||
len(source_symbols), parity_count
|
||
)
|
||
return tuple(
|
||
_linear_combine(row, source_symbols, symbol_size)
|
||
for row in generator[len(source_symbols):]
|
||
)
|
||
|
||
|
||
def outer_crc32(data: bytes) -> int:
|
||
return zlib.crc32(data) & 0xFFFFFFFF
|
||
|
||
|
||
def _validate_outer_symbol(symbol: OuterSymbol) -> None:
|
||
total_count = symbol.source_count + symbol.parity_count
|
||
if not 0 <= symbol.block_id <= 0xFFFFFFFF:
|
||
raise FECError("block_id is outside uint32")
|
||
if not 1 <= symbol.source_count <= MAX_BLOCK_SYMBOLS:
|
||
raise FECError("source_count is outside 1...255")
|
||
if not 0 <= symbol.parity_count <= (
|
||
MAX_BLOCK_SYMBOLS - symbol.source_count
|
||
):
|
||
raise FECError("invalid parity_count")
|
||
if not 0 <= symbol.symbol_index < total_count:
|
||
raise FECError("symbol_index is outside the block")
|
||
if not 1 <= symbol.symbol_size <= MAX_SYMBOL_SIZE:
|
||
raise FECError("symbol_size is outside uint16")
|
||
if not 1 <= len(symbol.data) <= symbol.symbol_size:
|
||
raise FECError("data length is outside symbol_size")
|
||
expected_parity = symbol.symbol_index >= symbol.source_count
|
||
if symbol.is_parity != expected_parity:
|
||
raise FECError("parity flag conflicts with symbol_index")
|
||
if symbol.is_parity and len(symbol.data) != symbol.symbol_size:
|
||
raise FECError("parity data must equal symbol_size")
|
||
|
||
|
||
def _pack_outer_header(
|
||
symbol: OuterSymbol,
|
||
crc_value: int,
|
||
) -> bytes:
|
||
return struct.pack(
|
||
OUTER_HEADER_FORMAT,
|
||
OUTER_MAGIC,
|
||
OUTER_VERSION,
|
||
OUTER_FLAG_PARITY if symbol.is_parity else 0,
|
||
OUTER_HEADER_SIZE,
|
||
symbol.block_id,
|
||
symbol.symbol_index,
|
||
symbol.source_count,
|
||
symbol.parity_count,
|
||
symbol.symbol_size,
|
||
len(symbol.data),
|
||
0,
|
||
0,
|
||
crc_value,
|
||
)
|
||
|
||
|
||
def encode_outer_symbol(symbol: OuterSymbol) -> bytes:
|
||
"""Serialize one source/parity symbol and calculate outer CRC32."""
|
||
|
||
if not isinstance(symbol, OuterSymbol):
|
||
raise TypeError("symbol must be OuterSymbol")
|
||
_validate_outer_symbol(symbol)
|
||
header_without_crc = _pack_outer_header(symbol, 0)
|
||
crc_value = outer_crc32(header_without_crc + symbol.data)
|
||
return _pack_outer_header(symbol, crc_value) + symbol.data
|
||
|
||
|
||
def decode_outer_symbol(wire_packet: bytes) -> OuterSymbol:
|
||
"""Deserialize and validate one outer FEC packet."""
|
||
|
||
if not isinstance(wire_packet, (bytes, bytearray)):
|
||
raise TypeError("wire_packet must be bytes or bytearray")
|
||
wire_packet = bytes(wire_packet)
|
||
if len(wire_packet) < OUTER_HEADER_SIZE + 1:
|
||
raise FECError("outer packet is too short")
|
||
(
|
||
magic,
|
||
version,
|
||
flags,
|
||
header_size,
|
||
block_id,
|
||
symbol_index,
|
||
source_count,
|
||
parity_count,
|
||
symbol_size,
|
||
data_length,
|
||
reserved16,
|
||
reserved32,
|
||
received_crc,
|
||
) = struct.unpack(
|
||
OUTER_HEADER_FORMAT,
|
||
wire_packet[:OUTER_HEADER_SIZE],
|
||
)
|
||
if magic != OUTER_MAGIC:
|
||
raise FECError("invalid outer magic")
|
||
if version != OUTER_VERSION:
|
||
raise FECError("unsupported outer version")
|
||
if flags & ~OUTER_FLAG_PARITY:
|
||
raise FECError("unsupported outer flags")
|
||
if header_size != OUTER_HEADER_SIZE:
|
||
raise FECError("invalid outer header size")
|
||
if reserved16 != 0 or reserved32 != 0:
|
||
raise FECError("reserved outer fields must be zero")
|
||
if len(wire_packet) != OUTER_HEADER_SIZE + data_length:
|
||
raise FECError("outer data_length does not match packet length")
|
||
|
||
symbol = OuterSymbol(
|
||
block_id=block_id,
|
||
symbol_index=symbol_index,
|
||
source_count=source_count,
|
||
parity_count=parity_count,
|
||
symbol_size=symbol_size,
|
||
data=wire_packet[OUTER_HEADER_SIZE:],
|
||
is_parity=bool(flags & OUTER_FLAG_PARITY),
|
||
outer_crc32=received_crc,
|
||
)
|
||
_validate_outer_symbol(symbol)
|
||
calculated_crc = outer_crc32(
|
||
_pack_outer_header(symbol, 0) + symbol.data
|
||
)
|
||
if calculated_crc != received_crc:
|
||
raise OuterPacketCRCError(
|
||
"outer CRC mismatch: "
|
||
f"received 0x{received_crc:08X}, "
|
||
f"calculated 0x{calculated_crc:08X}"
|
||
)
|
||
return symbol
|
||
|
||
|
||
def encode_fec_block(
|
||
inner_packets: tuple[bytes, ...],
|
||
block_id: int,
|
||
parity_count: int,
|
||
) -> tuple[bytes, ...]:
|
||
"""Wrap systematic Lab028 packets and append parity packets."""
|
||
|
||
if not inner_packets:
|
||
raise ValueError("inner packet block is empty")
|
||
if len(inner_packets) + parity_count > MAX_BLOCK_SYMBOLS:
|
||
raise ValueError("FEC block has more than 255 symbols")
|
||
validated_packets = []
|
||
for packet in inner_packets:
|
||
packet_bytes = bytes(packet)
|
||
decode_inner_packet(packet_bytes)
|
||
validated_packets.append(packet_bytes)
|
||
symbol_size = max(len(packet) for packet in validated_packets)
|
||
if symbol_size > MAX_SYMBOL_SIZE:
|
||
raise ValueError("inner packet is too large for outer symbol_size")
|
||
padded_sources = tuple(
|
||
packet + bytes(symbol_size - len(packet))
|
||
for packet in validated_packets
|
||
)
|
||
parity_symbols = encode_parity_symbols(
|
||
padded_sources, parity_count
|
||
)
|
||
source_count = len(validated_packets)
|
||
outer_packets = [
|
||
encode_outer_symbol(
|
||
OuterSymbol(
|
||
block_id=block_id,
|
||
symbol_index=index,
|
||
source_count=source_count,
|
||
parity_count=parity_count,
|
||
symbol_size=symbol_size,
|
||
data=packet,
|
||
is_parity=False,
|
||
)
|
||
)
|
||
for index, packet in enumerate(validated_packets)
|
||
]
|
||
outer_packets.extend(
|
||
encode_outer_symbol(
|
||
OuterSymbol(
|
||
block_id=block_id,
|
||
symbol_index=source_count + parity_index,
|
||
source_count=source_count,
|
||
parity_count=parity_count,
|
||
symbol_size=symbol_size,
|
||
data=parity,
|
||
is_parity=True,
|
||
)
|
||
)
|
||
for parity_index, parity in enumerate(parity_symbols)
|
||
)
|
||
return tuple(outer_packets)
|
||
|
||
|
||
def _padded_symbol(symbol: OuterSymbol) -> bytes:
|
||
return symbol.data + bytes(symbol.symbol_size - len(symbol.data))
|
||
|
||
|
||
def trim_recovered_inner_packet(padded_packet: bytes) -> bytes:
|
||
"""Trim GF padding using the recovered Lab028 header and validate CRC."""
|
||
|
||
if len(padded_packet) < INNER_HEADER_SIZE:
|
||
raise FECError("recovered symbol is shorter than Lab028 header")
|
||
unpacked = struct.unpack(
|
||
INNER_HEADER_FORMAT,
|
||
padded_packet[:INNER_HEADER_SIZE],
|
||
)
|
||
magic = unpacked[0]
|
||
version = unpacked[1]
|
||
payload_length = unpacked[7]
|
||
if magic != INNER_MAGIC or version != INNER_VERSION:
|
||
raise FECError("recovered Lab028 magic/version is invalid")
|
||
packet_length = INNER_HEADER_SIZE + payload_length
|
||
if packet_length > len(padded_packet):
|
||
raise FECError("recovered Lab028 length exceeds symbol_size")
|
||
inner_packet = padded_packet[:packet_length]
|
||
decode_inner_packet(inner_packet)
|
||
return inner_packet
|
||
|
||
|
||
def decode_fec_block(
|
||
wire_symbols: tuple[bytes, ...],
|
||
) -> DecodedFECBlock:
|
||
"""Recover all source Lab028 packets from any k valid outer symbols."""
|
||
|
||
if not wire_symbols:
|
||
raise InsufficientSymbolsError("no outer symbols received")
|
||
symbols_by_index: dict[int, OuterSymbol] = {}
|
||
metadata = None
|
||
for wire_symbol in wire_symbols:
|
||
symbol = decode_outer_symbol(wire_symbol)
|
||
current_metadata = (
|
||
symbol.block_id,
|
||
symbol.source_count,
|
||
symbol.parity_count,
|
||
symbol.symbol_size,
|
||
)
|
||
if metadata is None:
|
||
metadata = current_metadata
|
||
elif current_metadata != metadata:
|
||
raise FECError("outer symbols belong to different blocks")
|
||
existing = symbols_by_index.get(symbol.symbol_index)
|
||
if existing is not None:
|
||
if existing != symbol:
|
||
raise FECError("conflicting duplicate outer symbol")
|
||
continue
|
||
symbols_by_index[symbol.symbol_index] = symbol
|
||
|
||
assert metadata is not None
|
||
block_id, source_count, parity_count, symbol_size = metadata
|
||
if len(symbols_by_index) < source_count:
|
||
raise InsufficientSymbolsError(
|
||
f"need {source_count} symbols, got {len(symbols_by_index)}"
|
||
)
|
||
|
||
selected_indices = tuple(sorted(symbols_by_index)[:source_count])
|
||
generator = systematic_generator_matrix(
|
||
source_count, parity_count
|
||
)
|
||
decode_matrix = matrix_inverse(
|
||
tuple(generator[index] for index in selected_indices)
|
||
)
|
||
selected_symbols = tuple(
|
||
_padded_symbol(symbols_by_index[index])
|
||
for index in selected_indices
|
||
)
|
||
|
||
source_packets = []
|
||
recovered_indices = []
|
||
for source_index in range(source_count):
|
||
received_source = symbols_by_index.get(source_index)
|
||
if received_source is not None:
|
||
inner_packet = received_source.data
|
||
decode_inner_packet(inner_packet)
|
||
else:
|
||
padded_packet = _linear_combine(
|
||
decode_matrix[source_index],
|
||
selected_symbols,
|
||
symbol_size,
|
||
)
|
||
inner_packet = trim_recovered_inner_packet(padded_packet)
|
||
recovered_indices.append(source_index)
|
||
source_packets.append(inner_packet)
|
||
return DecodedFECBlock(
|
||
block_id=block_id,
|
||
source_packets=tuple(source_packets),
|
||
recovered_indices=tuple(recovered_indices),
|
||
)
|