diff --git a/lib/fabrics/ofi/CMakeLists.txt b/lib/fabrics/ofi/CMakeLists.txt index fef0cecce..d4b09c272 100644 --- a/lib/fabrics/ofi/CMakeLists.txt +++ b/lib/fabrics/ofi/CMakeLists.txt @@ -62,7 +62,7 @@ target_sources(mxl-fabrics-objects src/internal/Region.cpp src/internal/RegisteredRegion.cpp src/internal/RemoteRegion.cpp - src/internal/GrainSlices.cpp + src/internal/SliceRange.cpp src/internal/Target.cpp src/internal/TargetInfo.cpp ) diff --git a/lib/fabrics/ofi/src/internal/DataLayout.cpp b/lib/fabrics/ofi/src/internal/DataLayout.cpp index 6cfc45cb4..3aa0cb87c 100644 --- a/lib/fabrics/ofi/src/internal/DataLayout.cpp +++ b/lib/fabrics/ofi/src/internal/DataLayout.cpp @@ -4,16 +4,44 @@ #include "DataLayout.hpp" #include +#include #include +#include #include "mxl/flowinfo.h" +#include "Exception.hpp" namespace mxl::lib::fabrics::ofi { - DataLayout DataLayout::fromDiscrete(std::array const& sliceSizes) noexcept + DataLayout DataLayout::fromDiscrete(std::array const& sliceSizes, std::uint16_t totalSlices) noexcept { - return DataLayout{DataLayout::Discrete{.sliceSizes = sliceSizes}}; + return DataLayout{ + DataLayout::Discrete{.sliceSizes = sliceSizes, .totalSlices = totalSlices} + }; }; + std::size_t DataLayout::Discrete::totalLength() const noexcept + { + return std::accumulate(sliceSizes.begin(), sliceSizes.end(), std::size_t{0}); + } + + std::size_t DataLayout::Discrete::activePlaneCount() const noexcept + { + return std::ranges::count_if(sliceSizes, [](std::uint32_t sliceSize) { return sliceSize > 0; }); + } + + std::uint32_t DataLayout::Discrete::planePayloadOffset(std::size_t planeIndex, std::uint32_t grainPayloadOffset) const + { + if (sliceSizes.size() <= planeIndex) + { + throw Exception::invalidState("Invalid plane index {} not in range {}-{}", planeIndex, 0, sliceSizes.size()); + } + + return std::accumulate(sliceSizes.begin(), + sliceSizes.begin() + planeIndex, + grainPayloadOffset, + [this](std::uint32_t lhs, std::uint32_t rhs) { return lhs + (rhs * static_cast(totalSlices)); }); + } + DataLayout DataLayout::fromContinuous(std::size_t sampleSize, std::size_t channelCount, std::size_t bufferLength) noexcept { return DataLayout{ diff --git a/lib/fabrics/ofi/src/internal/DataLayout.hpp b/lib/fabrics/ofi/src/internal/DataLayout.hpp index 28e622a53..37f5c1e21 100644 --- a/lib/fabrics/ofi/src/internal/DataLayout.hpp +++ b/lib/fabrics/ofi/src/internal/DataLayout.hpp @@ -22,7 +22,24 @@ namespace mxl::lib::fabrics::ofi */ struct Discrete { - std::array sliceSizes; /**< Number of slices per plane. \see MXL_MAX_PLANES_PER_GRAIN */ + std::array + sliceSizes; /**< Size in bytes of a single slice for each plane. \see MXL_MAX_PLANES_PER_GRAIN */ + std::uint16_t totalSlices; /**< Total number of slices (e.g. video lines) per grain. */ + + /** \brief Return the total length of all active planes together */ + [[nodiscard]] + std::size_t totalLength() const noexcept; + + /** \brief Return the number of active planes (consecutive non-zero sliceSizes entries). */ + [[nodiscard]] + std::size_t activePlaneCount() const noexcept; + + /** \brief Return the byte offset from the grain start to where a given plane's payload begins. + * \param planeIndex The zero-based plane index. + * \param grainPayloadOffset The byte offset of the first plane's payload from the grain start (i.e. the grain header size). + */ + [[nodiscard]] + std::uint32_t planePayloadOffset(std::size_t planeIndex, std::uint32_t grainPayloadOffset) const; }; /** \brief Continuous layout variant of DataLayout. @@ -37,10 +54,12 @@ namespace mxl::lib::fabrics::ofi public: /** \brief Create a DataLayout representing video data. * \param sliceSizes The slice sizes of each planes in the video data layout. \see MXL_MAX_PLANES_PER_GRAIN + * \param totalSlices Total number of slices (e.g. video lines) per grain. * \return A DataLayout representing the specified video layout. */ [[nodiscard]] - static DataLayout fromDiscrete(std::array const& sliceSizes) noexcept; // NOLINT + static DataLayout fromDiscrete(std::array const& sliceSizes, + std::uint16_t totalSlices) noexcept; // NOLINT /** \brief Create a DataLayout representing audio data. * \param sampleSize The size of each audio sample in bytes. diff --git a/lib/fabrics/ofi/src/internal/Protocol.hpp b/lib/fabrics/ofi/src/internal/Protocol.hpp index 21a999da4..c9a880cf8 100644 --- a/lib/fabrics/ofi/src/internal/Protocol.hpp +++ b/lib/fabrics/ofi/src/internal/Protocol.hpp @@ -10,8 +10,8 @@ #include #include "DataLayout.hpp" #include "Endpoint.hpp" -#include "GrainSlices.hpp" #include "Region.hpp" +#include "SliceRange.hpp" #include "Target.hpp" #include "TargetInfo.hpp" diff --git a/lib/fabrics/ofi/src/internal/ProtocolEgressRMA.cpp b/lib/fabrics/ofi/src/internal/ProtocolEgressRMA.cpp index be5f5fc71..0dd77853c 100644 --- a/lib/fabrics/ofi/src/internal/ProtocolEgressRMA.cpp +++ b/lib/fabrics/ofi/src/internal/ProtocolEgressRMA.cpp @@ -28,16 +28,53 @@ namespace mxl::lib::fabrics::ofi void RMAGrainEgressProtocol::transferGrain(Endpoint const& ep, std::uint64_t localIndex, std::uint64_t remoteIndex, std::uint32_t payloadOffset, SliceRange const& sliceRange, ::fi_addr_t destAddr) { - auto const localSize = sliceRange.transferSize(payloadOffset, _layout.sliceSizes[0]); - auto const localOffset = sliceRange.transferOffset(payloadOffset, _layout.sliceSizes[0]); - auto const remoteSize = sliceRange.transferSize(payloadOffset, _layout.sliceSizes[0]); - auto const remoteOffset = sliceRange.transferOffset(payloadOffset, _layout.sliceSizes[0]); - - auto const localRegion = _localRegions[localIndex % _localRegions.size()].sub(localOffset, localSize); - auto const remoteRegion = _remoteInfo.remoteRegions[remoteIndex % _remoteInfo.remoteRegions.size()].sub(remoteOffset, remoteSize); + auto const localGrain = _localRegions[localIndex % _localRegions.size()]; + auto const remoteGrain = _remoteInfo.remoteRegions[remoteIndex % _remoteInfo.remoteRegions.size()]; auto const remoteSlot = remoteIndex % _remoteInfo.remoteRegions.size(); - _pending += ep.write(_token, localRegion, remoteRegion, destAddr, ImmDataGrain{remoteSlot, sliceRange.end()}.data()); + // Fast path: a full grain is transferred, can do one write no matter how many planes + if ((sliceRange.start() == 0) && (sliceRange.end() == _layout.totalSlices)) + { + _pending += ep.write(_token, localGrain, remoteGrain, destAddr, std::make_optional(ImmDataGrain{remoteSlot, _layout.totalSlices}.data())); + return; + } + + auto const planeCount = _layout.activePlaneCount(); + for (std::size_t plane = 0; plane < planeCount; ++plane) + { + auto const sliceSize = _layout.sliceSizes[plane]; + auto offset = std::uint32_t{0}; + auto size = std::uint32_t{0}; + + if (plane == 0) + { + // need to include header with slice 0-x + if (sliceRange.start() == 0) + { + offset = sliceRange.transferOffset(0, sliceSize); + size = sliceRange.transferSize(sliceSize) + payloadOffset; + } + else + { + offset = sliceRange.transferOffset(payloadOffset, sliceSize); + size = sliceRange.transferSize(sliceSize); + } + } + else + { + auto const planeBase = _layout.planePayloadOffset(plane, payloadOffset); + offset = sliceRange.transferOffset(planeBase, sliceSize); + size = sliceRange.transferSize(sliceSize); + } + + auto const localRegion = localGrain.sub(offset, size); + auto const remoteRegion = remoteGrain.sub(offset, size); + + auto const isLastPlane = (plane == planeCount - 1); + auto const immData = isLastPlane ? std::make_optional(ImmDataGrain{remoteSlot, sliceRange.end()}.data()) : std::nullopt; + + _pending += ep.write(_token, localRegion, remoteRegion, destAddr, immData); + } } void RMAGrainEgressProtocol::transferSamples(Endpoint const&, std::uint64_t, std::size_t, ::fi_addr_t) diff --git a/lib/fabrics/ofi/src/internal/ProtocolIngressRMA.cpp b/lib/fabrics/ofi/src/internal/ProtocolIngressRMA.cpp index e16a46a72..8bd26ab47 100644 --- a/lib/fabrics/ofi/src/internal/ProtocolIngressRMA.cpp +++ b/lib/fabrics/ofi/src/internal/ProtocolIngressRMA.cpp @@ -60,7 +60,7 @@ namespace mxl::lib::fabrics::ofi auto immData = completionData->data(); if (!immData) { - throw Exception::invalidState("Received a completion without immediate data."); + return {}; } auto [slot, slice] = ImmDataGrain{static_cast(*immData)}.unpack(); diff --git a/lib/fabrics/ofi/src/internal/RCInitiator.cpp b/lib/fabrics/ofi/src/internal/RCInitiator.cpp index 00b3c297f..6913e5fbd 100644 --- a/lib/fabrics/ofi/src/internal/RCInitiator.cpp +++ b/lib/fabrics/ofi/src/internal/RCInitiator.cpp @@ -16,9 +16,9 @@ #include "Exception.hpp" #include "FabricInfo.hpp" #include "FabricInfoHelpers.hpp" -#include "GrainSlices.hpp" #include "Protocol.hpp" #include "Region.hpp" +#include "SliceRange.hpp" #include "VariantUtils.hpp" namespace mxl::lib::fabrics::ofi diff --git a/lib/fabrics/ofi/src/internal/RDMInitiator.hpp b/lib/fabrics/ofi/src/internal/RDMInitiator.hpp index ebaf5b580..37bf01641 100644 --- a/lib/fabrics/ofi/src/internal/RDMInitiator.hpp +++ b/lib/fabrics/ofi/src/internal/RDMInitiator.hpp @@ -11,9 +11,9 @@ #include #include "mxl/fabrics.h" #include "Endpoint.hpp" -#include "GrainSlices.hpp" #include "Initiator.hpp" #include "Protocol.hpp" +#include "SliceRange.hpp" #include "TargetInfo.hpp" namespace mxl::lib::fabrics::ofi diff --git a/lib/fabrics/ofi/src/internal/Region.cpp b/lib/fabrics/ofi/src/internal/Region.cpp index 19589fad4..ef7045052 100644 --- a/lib/fabrics/ofi/src/internal/Region.cpp +++ b/lib/fabrics/ofi/src/internal/Region.cpp @@ -178,8 +178,9 @@ namespace mxl::lib::fabrics::ofi regions.emplace_back(grainInfoBaseAddr, grainInfoSize + grainPayloadSize, nullptr, nullptr, Region::Location::host()); } + auto const totalSlices = discreteFlow.grainAt(0)->header.info.totalSlices; return {std::move(regions), - DataLayout::fromDiscrete(std::to_array(discreteFlow.flowInfo()->config.discrete.sliceSizes)), + DataLayout::fromDiscrete(std::to_array(discreteFlow.flowInfo()->config.discrete.sliceSizes), totalSlices), discreteFlow.flowInfo()->config.common.maxSyncBatchSizeHint}; } else if (mxlIsContinuousDataFormat(static_cast(flow.flowInfo()->config.common.format))) diff --git a/lib/fabrics/ofi/src/internal/GrainSlices.cpp b/lib/fabrics/ofi/src/internal/SliceRange.cpp similarity index 63% rename from lib/fabrics/ofi/src/internal/GrainSlices.cpp rename to lib/fabrics/ofi/src/internal/SliceRange.cpp index e45d8b1a1..385ad9a31 100644 --- a/lib/fabrics/ofi/src/internal/GrainSlices.cpp +++ b/lib/fabrics/ofi/src/internal/SliceRange.cpp @@ -2,7 +2,7 @@ // // SPDX-License-Identifier: Apache-2.0 -#include "GrainSlices.hpp" +#include "SliceRange.hpp" #include "Exception.hpp" namespace mxl::lib::fabrics::ofi @@ -17,28 +17,14 @@ namespace mxl::lib::fabrics::ofi return SliceRange{start, end}; } - std::uint32_t SliceRange::transferSize(std::uint32_t payloadOffset, std::uint32_t sliceSize) const noexcept + std::uint32_t SliceRange::transferSize(std::uint32_t sliceSize) const noexcept { - auto size = (_end - _start) * sliceSize; - - if (_start == 0) - { - size += payloadOffset; - } - - return size; + return ((_end - _start) * sliceSize); } std::uint32_t SliceRange::transferOffset(std::uint32_t payloadOffset, std::uint32_t sliceSize) const noexcept { - if (_start == 0) - { - return 0; - } - else - { - return payloadOffset + (_start * sliceSize); - } + return payloadOffset + (_start * sliceSize); } std::uint16_t SliceRange::start() const noexcept diff --git a/lib/fabrics/ofi/src/internal/GrainSlices.hpp b/lib/fabrics/ofi/src/internal/SliceRange.hpp similarity index 94% rename from lib/fabrics/ofi/src/internal/GrainSlices.hpp rename to lib/fabrics/ofi/src/internal/SliceRange.hpp index bbc91918a..7b5fd23ad 100644 --- a/lib/fabrics/ofi/src/internal/GrainSlices.hpp +++ b/lib/fabrics/ofi/src/internal/SliceRange.hpp @@ -40,7 +40,7 @@ namespace mxl::lib::fabrics::ofi * \note When start is 0, the size includes the header, it adds the payload offset. */ [[nodiscard]] - std::uint32_t transferSize(std::uint32_t payloadOffset, std::uint32_t sliceSize) const noexcept; + std::uint32_t transferSize(std::uint32_t sliceSize) const noexcept; /** \brief Get the offset within the payload for the start of the range. * \note When start is 0, the offset is 0, because we include the header in the transfer. diff --git a/lib/tests/fabrics/CMakeLists.txt b/lib/tests/fabrics/CMakeLists.txt index a4e56d462..36aa042d9 100644 --- a/lib/tests/fabrics/CMakeLists.txt +++ b/lib/tests/fabrics/CMakeLists.txt @@ -26,6 +26,7 @@ target_sources(mxl-fabrics-tests PRIVATE test_basics.cpp test_interfaces.cpp + test_sliced.cpp ) target_link_libraries(mxl-fabrics-tests diff --git a/lib/tests/fabrics/ofi/CMakeLists.txt b/lib/tests/fabrics/ofi/CMakeLists.txt index 630602dca..15d464edd 100644 --- a/lib/tests/fabrics/ofi/CMakeLists.txt +++ b/lib/tests/fabrics/ofi/CMakeLists.txt @@ -31,6 +31,8 @@ target_sources(mxl-fabrics-ofi-tests test_Provider.cpp test_ProviderConfig.cpp test_Region.cpp + test_DataLayout.cpp + test_SliceRange.cpp ) target_link_libraries(mxl-fabrics-ofi-tests diff --git a/lib/tests/fabrics/ofi/Util.hpp b/lib/tests/fabrics/ofi/Util.hpp index 29d259c1a..d517e2177 100644 --- a/lib/tests/fabrics/ofi/Util.hpp +++ b/lib/tests/fabrics/ofi/Util.hpp @@ -77,7 +77,7 @@ namespace mxl::lib::fabrics::ofi inline MxlRegions getEmptyVideoMxlRegions() { - return MxlRegions({}, DataLayout::fromDiscrete({8, 0, 0, 0})); + return MxlRegions({}, DataLayout::fromDiscrete({8, 0, 0, 0}, 1)); } inline std::pair getHostRegionGroups() @@ -97,12 +97,12 @@ namespace mxl::lib::fabrics::ofi regions.emplace_back(*innerRegion.data(), innerRegion.size(), nullptr, nullptr); } - auto mxlRegions = MxlRegions(regions, DataLayout::fromDiscrete({8, 0, 0, 0})); + auto mxlRegions = MxlRegions(regions, DataLayout::fromDiscrete({8, 0, 0, 0}, 1)); return {mxlRegions, innerRegions}; } inline MxlRegions getMxlRegions(std::vector> const& innerRegions, - DataLayout dataLayout = DataLayout::fromDiscrete({8, 0, 0, 0})) + DataLayout dataLayout = DataLayout::fromDiscrete({8, 0, 0, 0}, 1)) { auto regions = std::vector{}; regions.reserve(innerRegions.size()); diff --git a/lib/tests/fabrics/ofi/test_DataLayout.cpp b/lib/tests/fabrics/ofi/test_DataLayout.cpp new file mode 100644 index 000000000..6aebc7893 --- /dev/null +++ b/lib/tests/fabrics/ofi/test_DataLayout.cpp @@ -0,0 +1,104 @@ +// SPDX-FileCopyrightText: 2026 Contributors to the Media eXchange Layer project. +// +// SPDX-License-Identifier: Apache-2.0 + +#include +#include "DataLayout.hpp" + +using namespace mxl::lib::fabrics::ofi; + +TEST_CASE("ofi: DataLayout::Discrete activePlaneCount", "[ofi][DataLayout]") +{ + SECTION("all zeros yields 0 planes") + { + auto layout = DataLayout::Discrete{ + .sliceSizes = {0, 0, 0, 0}, + .totalSlices = 0 + }; + REQUIRE(layout.activePlaneCount() == 0); + } + + SECTION("single plane (v210)") + { + auto layout = DataLayout::Discrete{ + .sliceSizes = {5120, 0, 0, 0}, + .totalSlices = 1080 + }; + REQUIRE(layout.activePlaneCount() == 1); + } + + SECTION("two planes (v210a)") + { + auto layout = DataLayout::Discrete{ + .sliceSizes = {5120, 2560, 0, 0}, + .totalSlices = 1080 + }; + REQUIRE(layout.activePlaneCount() == 2); + } + + SECTION("four planes") + { + auto layout = DataLayout::Discrete{ + .sliceSizes = {100, 200, 300, 400}, + .totalSlices = 10 + }; + REQUIRE(layout.activePlaneCount() == 4); + } +} + +TEST_CASE("ofi: DataLayout::Discrete planePayloadOffset", "[ofi][DataLayout]") +{ + constexpr auto headerSize = std::uint32_t{8192}; + + SECTION("single plane: plane 0 starts at the grain payload offset") + { + auto layout = DataLayout::Discrete{ + .sliceSizes = {5120, 0, 0, 0}, + .totalSlices = 1080 + }; + REQUIRE(layout.planePayloadOffset(0, headerSize) == headerSize); + } + + SECTION("two planes: plane 0 at header, plane 1 after fill data") + { + constexpr auto height = std::uint16_t{1080}; + constexpr auto fillSlice = std::uint32_t{5120}; + constexpr auto keySlice = std::uint32_t{2560}; + + auto layout = DataLayout::Discrete{ + .sliceSizes = {fillSlice, keySlice, 0, 0}, + .totalSlices = height + }; + + REQUIRE(layout.planePayloadOffset(0, headerSize) == headerSize); + REQUIRE(layout.planePayloadOffset(1, headerSize) == headerSize + (height * fillSlice)); + } + + SECTION("four planes: each plane offset accumulates") + { + constexpr auto slices = std::uint32_t{10}; + auto layout = DataLayout::Discrete{ + .sliceSizes = {100, 200, 300, 400}, + .totalSlices = slices + }; + + REQUIRE(layout.planePayloadOffset(0, headerSize) == headerSize); + REQUIRE(layout.planePayloadOffset(1, headerSize) == headerSize + (slices * 100)); + REQUIRE(layout.planePayloadOffset(2, headerSize) == headerSize + (slices * 100) + (slices * 200)); + REQUIRE(layout.planePayloadOffset(3, headerSize) == headerSize + (slices * 100) + (slices * 200) + (slices * 300)); + } +} + +TEST_CASE("ofi: DataLayout fromDiscrete factory", "[ofi][DataLayout]") +{ + auto sliceSizes = std::array{5120, 2560, 0, 0}; + auto layout = DataLayout::fromDiscrete(sliceSizes, 1080); + + REQUIRE(layout.isDiscrete()); + REQUIRE_FALSE(layout.isContinuous()); + + auto const& discrete = layout.asDiscrete(); + REQUIRE(discrete.sliceSizes == sliceSizes); + REQUIRE(discrete.totalSlices == 1080); + REQUIRE(discrete.activePlaneCount() == 2); +} diff --git a/lib/tests/fabrics/ofi/test_SliceRange.cpp b/lib/tests/fabrics/ofi/test_SliceRange.cpp new file mode 100644 index 000000000..e8f87f7b9 --- /dev/null +++ b/lib/tests/fabrics/ofi/test_SliceRange.cpp @@ -0,0 +1,60 @@ +// SPDX-FileCopyrightText: 2026 Contributors to the Media eXchange Layer project. +// +// SPDX-License-Identifier: Apache-2.0 + +#include +#include "SliceRange.hpp" + +using namespace mxl::lib::fabrics::ofi; + +TEST_CASE("ofi: SliceRange construction", "[ofi][GrainSlices]") +{ + SECTION("valid range") + { + auto range = SliceRange::make(10, 20); + REQUIRE(range.start() == 10); + REQUIRE(range.end() == 20); + } + + SECTION("single-element range") + { + auto range = SliceRange::make(5, 6); + REQUIRE(range.start() == 5); + REQUIRE(range.end() == 6); + } + + SECTION("empty range (start == end) is valid") + { + auto range = SliceRange::make(0, 0); + REQUIRE(range.start() == 0); + REQUIRE(range.end() == 0); + } + + SECTION("inverted range throws") + { + REQUIRE_THROWS(SliceRange::make(10, 5)); + } +} + +TEST_CASE("ofi: SliceRange transferSize", "[ofi][GrainSlices]") +{ + constexpr auto const sliceSize = std::uint32_t{5120}; + + SECTION("from nonzero start excludes header") + { + auto range = SliceRange::make(540, 1080); + REQUIRE(range.transferSize(sliceSize) == 540 * sliceSize); + } +} + +TEST_CASE("ofi: SliceRange transferOffset", "[ofi][GrainSlices]") +{ + constexpr auto const payloadOffset = std::uint32_t{8192}; + constexpr auto const sliceSize = std::uint32_t{5120}; + + SECTION("from nonzero start returns offset past header") + { + auto range = SliceRange::make(540, 1080); + REQUIRE(range.transferOffset(payloadOffset, sliceSize) == payloadOffset + (540 * sliceSize)); + } +} diff --git a/lib/tests/fabrics/test_sliced.cpp b/lib/tests/fabrics/test_sliced.cpp new file mode 100644 index 000000000..1a63e20fd --- /dev/null +++ b/lib/tests/fabrics/test_sliced.cpp @@ -0,0 +1,547 @@ +// SPDX-FileCopyrightText: 2026 Contributors to the Media eXchange Layer project. +// +// SPDX-License-Identifier: Apache-2.0 + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include "mxl/flow.h" + +namespace +{ + constexpr auto TEST_OP_TIMEOUT = std::chrono::seconds(5); + + template + void check(F f, std::string const& msg, Args&&... args) + { + auto const ret = f(std::forward(args)...); + if (ret != MXL_STATUS_OK) + { + throw std::runtime_error{msg + " (status " + std::to_string(ret) + ")"}; + } + } + + std::string readFile(std::filesystem::path const& filepath) + { + if (auto file = std::ifstream{filepath, std::ios::in | std::ios::binary}; file) + { + auto reader = std::stringstream{}; + reader << file.rdbuf(); + return reader.str(); + } + throw std::runtime_error("Failed to open file: " + filepath.string()); + } + + struct ProviderTCP + { + constexpr static auto provider = MXL_FABRICS_PROVIDER_TCP; + constexpr static char const* targetNode = "127.0.0.1"; + constexpr static char const* targetService = "0"; + constexpr static char const* initiatorNode = "127.0.0.1"; + constexpr static char const* initiatorService = "0"; + }; + + struct ProviderSHM + { + constexpr static auto provider = MXL_FABRICS_PROVIDER_SHM; + constexpr static char const* targetNode = "target"; + constexpr static char const* targetService = "sliced"; + constexpr static char const* initiatorNode = "initiator"; + constexpr static char const* initiatorService = "sliced"; + }; + + struct FlowV210 + { + constexpr static char const* flowDefPath = "../data/v210_flow.json"; + constexpr static char const* flowId = "5fbec3b1-1b0f-417d-9059-8b94a47197ed"; + constexpr static auto hasAlphaChannel = false; + }; + + struct FlowV210a + { + constexpr static char const* flowDefPath = "../data/v210a_flow.json"; + constexpr static char const* flowId = "5fbec3b1-1b0f-417d-9059-8b94a47197ed"; + constexpr static auto hasAlphaChannel = true; + }; + + template + struct TestVariant + : Provider + , Flow + {}; + + using TCP_V210 = TestVariant; + using TCP_V210a = TestVariant; + using SHM_V210 = TestVariant; + using SHM_V210a = TestVariant; + + class TempDomainGuard + { + public: + TempDomainGuard() + : _path() + { + char path[] = "/dev/shm/mxl-domain-XXXXXX"; + if (::mkdtemp(path) == nullptr) + { + throw std::runtime_error{"failed to create temporary domain"}; + } + + _path = path; + } + + ~TempDomainGuard() + { + std::filesystem::remove_all(_path); + } + + char const* c_str() const + { + return _path.c_str(); + } + + private: + std::string _path; + }; + + template + class FabricsTestFixture + { + public: + FabricsTestFixture() + : _targetDomain() + , _initiatorDomain() + { + _initiatorInstance = mxlCreateInstance(_initiatorDomain.c_str(), nullptr); + if (_initiatorInstance == nullptr) + { + throw std::runtime_error{"failed to create initiator mxl instance"}; + } + _targetInstance = mxlCreateInstance(_targetDomain.c_str(), nullptr); + if (_targetInstance == nullptr) + { + throw std::runtime_error{"failed to create target mxl instance"}; + } + + auto const flowDef = readFile(ProviderType::flowDefPath); + auto const* flowId = ProviderType::flowId; + + // clang-format off + check(mxlFabricsCreateInstance, "failed to create initiator fabrics instance", + _initiatorInstance, nullptr, &_initiatorFabricsInstance); + check(mxlFabricsCreateInstance, "failed to create target fabrics instance", + _targetInstance, nullptr, &_targetFabricsInstance); + check(mxlCreateFlowWriter, "failed to create initiator flow writer", + _initiatorInstance, flowDef.c_str(), nullptr, &_initiatorWriter, &_flowConfigInfo, nullptr); + check(mxlCreateFlowReader, "failed to create initiator flow reader", + _initiatorInstance, flowId, nullptr, &_initiatorReader); + check(mxlCreateFlowWriter, "failed to create target flow writer", + _targetInstance, flowDef.c_str(), nullptr, &_targetWriter, nullptr, nullptr); + check(mxlCreateFlowReader, "failed to create target flow reader", + _targetInstance, flowId, nullptr, &_targetReader); + check(mxlFabricsCreateTarget, "failed to create fabrics target", + _targetFabricsInstance, &_target); + check(mxlFabricsCreateInitiator, "failed to create fabrics initiator", + _initiatorFabricsInstance, &_initiator); + + auto targetConfig = mxlFabricsTargetConfig{ + .version = MXL_FABRICS_API_VERSION, + .interface = { + .version = MXL_FABRICS_API_VERSION, + .provider = ProviderType::provider, + .caps = { .version = MXL_FABRICS_API_VERSION, .flags = MXL_FABRICS_IFACE_CAP_REMOTE_WRITE | MXL_FABRICS_IFACE_CAP_BLOCKING_OPERATIONS, .maxMessageSize = 0 }, + .address = {.node = ProviderType::targetNode, .service = ProviderType::targetService}, + .attr = nullptr, + }, + .writer = _targetWriter, + }; + check(mxlFabricsTargetSetup, "failed to set up target", + _target, &targetConfig, nullptr, &_targetInfo); + + auto initiatorConfig = mxlFabricsInitiatorConfig{ + .version = MXL_FABRICS_API_VERSION, + .interface = { + .version = MXL_FABRICS_API_VERSION, + .provider = ProviderType::provider, + .caps = { .version = MXL_FABRICS_API_VERSION, .flags = MXL_FABRICS_IFACE_CAP_REMOTE_WRITE | MXL_FABRICS_IFACE_CAP_BLOCKING_OPERATIONS , .maxMessageSize = 0 }, + .address = {.node = ProviderType::initiatorNode, .service = ProviderType::initiatorService}, + .attr = nullptr, + }, + .reader = _initiatorReader, + }; + + check(mxlFabricsInitiatorSetup, "failed to set up initiator", + _initiator, &initiatorConfig, nullptr); + check(mxlFabricsInitiatorAddTarget, "failed to add target to initiator", + _initiator, _targetInfo); + // clang-format on + + driveConnectionProgress(); + } + + virtual ~FabricsTestFixture() + { + mxlFabricsDestroyInitiator(_initiatorFabricsInstance, _initiator); + mxlFabricsDestroyTarget(_targetFabricsInstance, _target); + mxlFabricsFreeTargetInfo(_targetInfo); + + mxlReleaseFlowReader(_targetInstance, _targetReader); + mxlReleaseFlowWriter(_targetInstance, _targetWriter); + mxlReleaseFlowReader(_initiatorInstance, _initiatorReader); + mxlReleaseFlowWriter(_initiatorInstance, _initiatorWriter); + + mxlFabricsDestroyInstance(_initiatorFabricsInstance); + mxlFabricsDestroyInstance(_targetFabricsInstance); + mxlDestroyInstance(_targetInstance); + mxlDestroyInstance(_initiatorInstance); + } + + protected: + bool hasAlphaChannel() const noexcept + { + return ProviderType::hasAlphaChannel; + } + + // Fill all planes of the initiator grain with zero + void fillInitiatorGrainWithZeros(std::uint64_t index) + { + fillGrainWithZeros(_initiatorWriter, index); + } + + // Fill all planes of the target grain with zero + void fillTargetGrainWithZeros(std::uint64_t index) + { + fillGrainWithZeros(_targetWriter, index); + } + + // Get the grain header on the targe side + mxlGrainInfo getTargetGrainInfo(std::uint64_t index) + { + auto payload = std::add_pointer_t{nullptr}; + auto grainInfo = mxlGrainInfo{}; + REQUIRE(mxlFlowReaderGetGrainNonBlocking(_targetReader, index, &grainInfo, &payload) == MXL_STATUS_OK); + return grainInfo; + } + + // Get the grain header on the initiator side + mxlGrainInfo getInitiatorGrainInfo(std::uint64_t index) + { + auto payload = std::add_pointer_t{nullptr}; + auto grainInfo = mxlGrainInfo{}; + REQUIRE(mxlFlowReaderGetGrainNonBlocking(_initiatorReader, index, &grainInfo, &payload) == MXL_STATUS_OK); + return grainInfo; + } + + // Write a slice of key and fill data on the initiator grain + void writeInitiatorGrainSlice(std::uint64_t index, std::uint32_t sliceIndex, std::uint8_t fill, std::uint8_t key) + { + auto grainInfo = mxlGrainInfo{}; + auto payload = static_cast(nullptr); + REQUIRE(mxlFlowWriterOpenGrain(_initiatorWriter, index, &grainInfo, &payload) == MXL_STATUS_OK); + + auto planeOffset = std::uint32_t{0}; + auto planeIndex = std::size_t{0}; + for (auto const sliceLen : _flowConfigInfo.discrete.sliceSizes) + { + if (sliceLen == 0) + { + break; + } + + auto sliceStart = payload + planeOffset + (static_cast(sliceLen) * sliceIndex); + std::memset(sliceStart, planeIndex == 0 ? fill : key, sliceLen); + planeOffset += static_cast(sliceLen) * grainInfo.totalSlices; + ++planeIndex; + } + + grainInfo.validSlices = static_cast(sliceIndex + 1); + REQUIRE(mxlFlowWriterCommitGrain(_initiatorWriter, &grainInfo) == MXL_STATUS_OK); + } + + // Run a full grain slice range transfer + std::uint64_t transferGrainSlices(std::uint64_t grainIndex, std::uint16_t startSlice, std::uint16_t endSlice) + { + auto deadline = std::chrono::steady_clock::now() + TEST_OP_TIMEOUT; + auto initiatorDone = false; + auto targetDone = false; + auto readIndex = std::uint64_t{0}; + auto status = MXL_STATUS_OK; + + // Enqueue the transfer until successfull + // The SHM provider needs the progress and read calls even if no transfers have been queued yet to become ready. + do + { + status = mxlFabricsInitiatorTransferGrain(_initiator, grainIndex, startSlice, endSlice); + mxlFabricsInitiatorMakeProgressNonBlocking(_initiator); // the SHM provider needs a CQ read here sometimes to not get stuck + if (mxlFabricsTargetReadGrainNonBlocking(_target, &readIndex) == MXL_STATUS_OK) + { + // The SHM provider is sometimes ready right after the write has been enqueued. + return readIndex; + } + } + while (status != MXL_STATUS_OK && std::chrono::steady_clock::now() < deadline); + + // Timeout of not ok + REQUIRE(status == MXL_STATUS_OK); + + // Drive progress on both sides of the connection until the grain slices have been transferred + while (std::chrono::steady_clock::now() < deadline) + { + if (!initiatorDone) + { + auto status = mxlFabricsInitiatorMakeProgressNonBlocking(_initiator); + if (status == MXL_STATUS_OK) + { + initiatorDone = true; + } + else + { + REQUIRE(status == MXL_ERR_NOT_READY); + } + } + + if (!targetDone) + { + auto status = mxlFabricsTargetReadGrainNonBlocking(_target, &readIndex); + if (status == MXL_STATUS_OK) + { + targetDone = true; + } + else + { + REQUIRE(status == MXL_ERR_NOT_READY); + } + } + + if (initiatorDone && targetDone) + { + return readIndex; + } + } + + FAIL("grain slice transfer did not complete within timeout"); + return readIndex; + } + + // Read the first byte of a slice in the key and fill buffer + std::pair readTargetGrainSlice(std::uint64_t index, std::uint32_t sliceIndex) + { + auto grainInfo = mxlGrainInfo{}; + auto payload = static_cast(nullptr); + REQUIRE(mxlFlowReaderGetGrainSliceNonBlocking(_targetReader, index, static_cast(sliceIndex + 1), &grainInfo, &payload) == + MXL_STATUS_OK); + + auto result = std::pair{0, 0}; + auto planeOffset = std::uint32_t{0}; + auto planeIndex = std::size_t{0}; + for (auto const sliceLen : _flowConfigInfo.discrete.sliceSizes) + { + if (sliceLen == 0) + { + break; + } + + auto byte = *(payload + (planeOffset + static_cast(sliceLen) * sliceIndex)); + if (planeIndex == 0) + { + result.first = byte; + } + else + { + result.second = byte; + } + + planeOffset += static_cast(sliceLen) * grainInfo.totalSlices; + ++planeIndex; + } + + return result; + } + + private: + // fill a grain at "index" with zero + void fillGrainWithZeros(mxlFlowWriter writer, std::uint64_t index) + { + auto grainInfo = mxlGrainInfo{}; + auto payload = static_cast(nullptr); + REQUIRE(mxlFlowWriterOpenGrain(writer, index, &grainInfo, &payload) == MXL_STATUS_OK); + + auto planeOffset = std::uint32_t{0}; + for (auto const sliceLen : _flowConfigInfo.discrete.sliceSizes) + { + if (sliceLen == 0) + { + break; + } + + std::memset(payload + planeOffset, 0, static_cast(sliceLen) * grainInfo.totalSlices); + planeOffset += static_cast(sliceLen) * grainInfo.totalSlices; + } + + grainInfo.validSlices = 0; + grainInfo.flags = MXL_GRAIN_FLAG_INVALID; + REQUIRE(mxlFlowWriterCommitGrain(writer, &grainInfo) == MXL_STATUS_OK); + } + + // drive progress for the connection to be established + void driveConnectionProgress() + { + auto deadline = std::chrono::steady_clock::now() + TEST_OP_TIMEOUT; + while (std::chrono::steady_clock::now() < deadline) + { + auto dummyIndex = std::uint64_t{0}; + mxlFabricsTargetReadGrainNonBlocking(_target, &dummyIndex); + + auto status = mxlFabricsInitiatorMakeProgressNonBlocking(_initiator); + if (status == MXL_STATUS_OK) + { + return; + } + if (status != MXL_ERR_NOT_READY) + { + throw std::runtime_error{"initiator progress failed with status " + std::to_string(status)}; + } + } + throw std::runtime_error{"failed to establish connection within timeout"}; + } + + private: + TempDomainGuard _targetDomain; + TempDomainGuard _initiatorDomain; + + mxlInstance _initiatorInstance = nullptr; + mxlInstance _targetInstance = nullptr; + mxlFabricsInstance _initiatorFabricsInstance = nullptr; + mxlFabricsInstance _targetFabricsInstance = nullptr; + mxlFabricsTargetInfo _targetInfo = nullptr; + + mxlFlowConfigInfo _flowConfigInfo = {}; + + mxlFlowWriter _initiatorWriter = nullptr; + mxlFlowReader _initiatorReader = nullptr; + mxlFlowWriter _targetWriter = nullptr; + mxlFlowReader _targetReader = nullptr; + + mxlFabricsTarget _target = nullptr; + mxlFabricsInitiator _initiator = nullptr; + }; +} + +TEMPLATE_TEST_CASE_METHOD(FabricsTestFixture, "Slice transfer single", "[sliced][single-slices]", TCP_V210, TCP_V210a /*, SHM_V210, SHM_V210a */) +{ + constexpr auto const startGrainIndex = std::uint64_t{140}; + for (auto iteration = std::size_t{0}; iteration < 4; ++iteration) + { + auto lastSlice = std::uint16_t{0}; + auto const grainIndex = startGrainIndex + iteration; + this->fillInitiatorGrainWithZeros(grainIndex); + this->fillTargetGrainWithZeros(grainIndex); + + auto info = this->getInitiatorGrainInfo(grainIndex); + for (auto slice = std::uint16_t{0}; slice < info.totalSlices; ++slice) + { + // check that the target slice is zeroed + auto [fillBefore, keyBefore] = this->readTargetGrainSlice(grainIndex, slice); + REQUIRE(fillBefore == 0x00); + if (this->hasAlphaChannel()) + { + REQUIRE(keyBefore == 0x00); + } + + auto fillValue = static_cast(0xAC & ~slice); + auto keyValue = static_cast(0xAB & ~slice); + + // write pattern to initiator grain buffer + this->writeInitiatorGrainSlice(grainIndex, slice, fillValue, keyValue); + REQUIRE(this->transferGrainSlices(grainIndex, lastSlice, slice + 1) == grainIndex); + this->readTargetGrainSlice(grainIndex, slice); + + // read after transfer + auto [fillAfter, keyAfter] = this->readTargetGrainSlice(grainIndex, slice); + REQUIRE(fillAfter == fillValue); + if (this->hasAlphaChannel()) + { + REQUIRE(keyAfter == keyValue); + } + + // check that validSlices has been updated + auto const targetGrainInfo = this->getTargetGrainInfo(grainIndex); + REQUIRE(targetGrainInfo.validSlices == slice + 1); + } + } +} + +TEMPLATE_TEST_CASE_METHOD(FabricsTestFixture, "Slice transfer blocks", "[sliced][slice-blocks]", TCP_V210, TCP_V210a /*, SHM_V210, SHM_V210a */) +{ + // test different block sizes + constexpr auto const blockSizes = std::array{2, 13, 73, 720}; + for (auto const blockSize : blockSizes) + { + auto const startGrainIndex = std::uint64_t{140 + static_cast(blockSize)}; + for (auto iteration = std::size_t{0}; iteration < 4; ++iteration) + { + auto lastSlice = std::uint16_t{0}; + auto const grainIndex = startGrainIndex + iteration; + this->fillInitiatorGrainWithZeros(grainIndex); + this->fillTargetGrainWithZeros(grainIndex); + + auto startSlice = std::uint16_t{0}; + auto info = this->getInitiatorGrainInfo(grainIndex); + for (;;) + { + auto endSlice = startSlice + blockSize; + if (endSlice >= info.totalSlices) + { + endSlice = info.totalSlices; + } + + // find a slice in the middle of the block <= (endSlice - 1) + auto middleSlice = std::min(startSlice + (blockSize / 2), endSlice - 1); + auto [fillValueBefore, keyValueBefore] = this->readTargetGrainSlice(grainIndex, middleSlice); + REQUIRE(fillValueBefore == 0x00); + if (this->hasAlphaChannel()) + { + REQUIRE(keyValueBefore == 0x00); + } + + // generate a value for fill and key buffer + auto fillValue = static_cast(0xAC & ~middleSlice); + auto keyValue = static_cast(0xAB & ~middleSlice); + + // write to a slice in the middle of the block + this->writeInitiatorGrainSlice(grainIndex, middleSlice, fillValue, keyValue); + + // transfer the slices + REQUIRE(this->transferGrainSlices(grainIndex, lastSlice, endSlice) == grainIndex); + + // check that the payload is correct + auto [fillValueAfter, keyValueAfter] = this->readTargetGrainSlice(grainIndex, middleSlice); + REQUIRE(fillValueAfter == fillValue); + if (this->hasAlphaChannel()) + { + REQUIRE(keyValueAfter == keyValue); + } + + // check that validSlices was updated + auto const targetGrainInfo = this->getTargetGrainInfo(grainIndex); + REQUIRE(targetGrainInfo.validSlices == endSlice); + + startSlice += blockSize; + if (startSlice >= info.totalSlices) + { + break; + } + } + } + } +}