Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions src/signaloid/distributional/distributional.py
Original file line number Diff line number Diff line change
Expand Up @@ -718,6 +718,9 @@ def export(self, to_str: bool = True) -> str | bytes:
UxString += "Ux"

fmt = STRUCT_FORMATS["str"]

# # Representation type (uint8_t) (1 byte)
buffer += struct.pack(fmt["UR_type"], self.UR_type)
else:
# Particle value (double) (8 bytes)
particle_value = (
Expand All @@ -733,8 +736,11 @@ def export(self, to_str: bool = True) -> str | bytes:
# representation type.
buffer += bytes([UX_BINARY_FORMAT_MARKER, 0x00, 0x00])

# Representation type (uint8_t) (1 byte)
buffer += struct.pack(fmt["UR_type"], self.UR_type)
# Representation type (uint8_t) (1 byte)
buffer += struct.pack(
fmt["UR_type"],
UR_TYPE_ATHENS if self.UR_type == UR_TYPE_UNSPECIFIED else self.UR_type,
)

# Number of samples (uint64_t) (8 bytes)
buffer += struct.pack(fmt["sample_count"], self.UR_order)
Expand Down
25 changes: 20 additions & 5 deletions src/signaloid/distributional/distributional_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,9 @@

import numpy as np
from signaloid.distributional.dirac_delta import DiracDelta
from signaloid.distributional.distributional import DistributionalValue
from signaloid.distributional.distributional import (
DistributionalValue,
)


def read_string_bytes_pairs_from_csv(
Expand Down Expand Up @@ -71,8 +73,19 @@ def to_padded_ux_binary(ux_binary: bytes) -> bytes:
Returns:
The equivalent Ux Binary Data array.
"""
if len(ux_binary) >= 9 and ux_binary[8:11] == b"\xf0\x00\x00":
if (
len(ux_binary) >= 12
and ux_binary[8:11] == b"\xf0\x00\x00"
and ux_binary[11] != 0x00
):
return ux_binary

if ux_binary[8] == 0x00:
return ux_binary[:8] + b"\xf0\x00\x00\x04" + ux_binary[9:]

if ux_binary[8:12] == b"\xf0\x00\x00\x00":
return ux_binary[:8] + b"\xf0\x00\x00\x04" + ux_binary[12:]

return ux_binary[:8] + b"\xf0\x00\x00" + ux_binary[8:]


Expand Down Expand Up @@ -1205,7 +1218,7 @@ class TestUxBinaryFormatDetection(unittest.TestCase):

def _padded_format_hex(self) -> str:
"""The LEGACY_FORMAT_HEX value rewritten in the Ux Binary layout."""
return self.LEGACY_FORMAT_HEX[:16] + "f00000" + self.LEGACY_FORMAT_HEX[16:]
return self.LEGACY_FORMAT_HEX[:16] + "f0000004" + self.LEGACY_FORMAT_HEX[18:]

def test_export_writes_format_marker(self) -> None:
"""`bytes(dist)` places 0xF0 0x00 0x00 right after the particle."""
Expand Down Expand Up @@ -1236,7 +1249,7 @@ def test_export_round_trips_through_parse(self) -> None:
self.assertIsNotNone(parsed)
assert parsed is not None
self.assertEqual(parsed.particle_value, dist.particle_value)
self.assertEqual(parsed.UR_type, dist.UR_type)
self.assertEqual(parsed.UR_type, dist.UR_type if dist.UR_type != 0x00 else 0x04)
np.testing.assert_array_equal(parsed.positions, dist.positions)
np.testing.assert_array_equal(parsed.raw_masses, dist.raw_masses)

Expand All @@ -1249,7 +1262,9 @@ def test_legacy_and_padded_layouts_parse_equal(self) -> None:
self.assertIsNotNone(padded)
assert legacy is not None and padded is not None
self.assertEqual(legacy.particle_value, padded.particle_value)
self.assertEqual(legacy.UR_type, padded.UR_type)
self.assertEqual(
legacy.UR_type if legacy.UR_type != 0x00 else 0x04, padded.UR_type
)
self.assertEqual(legacy.UR_order, padded.UR_order)
np.testing.assert_array_equal(legacy.positions, padded.positions)
np.testing.assert_array_equal(legacy.raw_masses, padded.raw_masses)
Expand Down
Loading