From 8e56859da9ce2f90b783b8d42e709980a61edb9f Mon Sep 17 00:00:00 2001 From: Xiao Wang Date: Sat, 15 Aug 2026 16:19:36 -0500 Subject: [PATCH 1/7] Add uint8_t support to bit-oriented APIs Accept byte-bool buffers in bit packing, PRG output, packed IOChannel transfers, F2k selectors, and Galois-field packing without casting them to bool pointers. Share implementations with the existing bool overloads. Preserve null-safe zero-length calls, normalize byte-bool output, and keep the bool and uint8_t paths bit-, wire-, and PRG-state-equivalent. --- bench/bench_block.cpp | 6 +-- bench/bench_f2k.cpp | 13 +++---- bench/bench_prg.cpp | 2 +- docs/api_conventions.md | 10 +++-- emp-tool/runtime/core/utils.h | 11 +++++- emp-tool/runtime/core/utils.hpp | 33 ++++++++++++++-- emp-tool/runtime/crypto/f2k.h | 11 ++++++ emp-tool/runtime/crypto/f2k.hpp | 42 +++++++++++++++++--- emp-tool/runtime/crypto/prg.h | 66 ++++++++++++++++++-------------- emp-tool/runtime/io/io_channel.h | 26 ++++++++++++- test/runtime/test_block.cpp | 34 ++++++++++------ test/runtime/test_f2k.cpp | 33 +++++++++------- test/runtime/test_netio.cpp | 45 +++++++++++++++++----- test/runtime/test_prg.cpp | 30 +++++++++++++-- 14 files changed, 269 insertions(+), 93 deletions(-) diff --git a/bench/bench_block.cpp b/bench/bench_block.cpp index 79150ad..2c21194 100644 --- a/bench/bench_block.cpp +++ b/bench/bench_block.cpp @@ -155,10 +155,10 @@ static void bench(double sec) { cout << "\n=== bools_to_bits / bits_to_bools (sweep N bits) ===\n"; for (int len : {32, 128, 1024, 8192, 65536}) { vector bools(len); - prg.random_bool(reinterpret_cast(bools.data()), len); + prg.random_bool(bools.data(), len); vector packed((len + 7) / 8); double calls = run_for(sec, [&]() { - bools_to_bits(packed.data(), reinterpret_cast(bools.data()), len); + bools_to_bits(packed.data(), bools.data(), len); }, packed.data()); ostringstream lbl; lbl << "bools_to_bits(N=" << len << ")"; // Bandwidth: bytes of bool input read. @@ -169,7 +169,7 @@ static void bench(double sec) { prg.random_data_unaligned(packed.data(), (int)packed.size()); vector bools(len); double calls = run_for(sec, [&]() { - bits_to_bools(reinterpret_cast(bools.data()), packed.data(), len); + bits_to_bools(bools.data(), packed.data(), len); }, bools.data()); ostringstream lbl; lbl << "bits_to_bools(N=" << len << ")"; print_vec_bytes(lbl.str(), calls, (size_t)len); diff --git a/bench/bench_f2k.cpp b/bench/bench_f2k.cpp index 9a765ed..a00482e 100644 --- a/bench/bench_f2k.cpp +++ b/bench/bench_f2k.cpp @@ -117,13 +117,12 @@ static void bench(double sec) { vector a(n); vector bs(n); prg.random_block(a.data(), n); - prg.random_bool(reinterpret_cast(bs.data()), n); + prg.random_bool(bs.data(), n); block r; double calls = run_for(sec, [&]() { - vector_inn_prdt_sum_red(&r, a.data(), - reinterpret_cast(bs.data()), n); + vector_inn_prdt_sum_red(&r, a.data(), bs.data(), n); }, &r); - ostringstream lbl; lbl << "vec_inn_prdt_red(bool, N=" << n << ")"; + ostringstream lbl; lbl << "vec_inn_prdt_red(byte-bool, N=" << n << ")"; print_vec(lbl.str(), calls, n); } @@ -174,12 +173,12 @@ static void bench(double sec) { { GaloisFieldPacking pkr; uint8_t bits[128]; - prg.random_bool(reinterpret_cast(bits), 128); + prg.random_bool(bits, 128); block r; double calls = run_for(sec, [&]() { - pkr.packing(&r, reinterpret_cast(bits)); + pkr.packing(&r, bits); }, &r); - print_op("packing(bool*, 128)", calls); + print_op("packing(byte-bool*, 128)", calls); } } diff --git a/bench/bench_prg.cpp b/bench/bench_prg.cpp index f475302..fcbbf87 100644 --- a/bench/bench_prg.cpp +++ b/bench/bench_prg.cpp @@ -89,7 +89,7 @@ static void bench(double sec) { for (int nb : {32, 128, 512, 2048, 8192, 32768, 131072}) { vector buf(nb); double calls = run_for(sec, [&]() { - prg.random_bool(reinterpret_cast(buf.data()), nb); + prg.random_bool(buf.data(), nb); }, buf.data()); ostringstream lbl; lbl << "random_bool(N=" << nb << ")"; print_vec(lbl.str(), calls, (size_t)nb); diff --git a/docs/api_conventions.md b/docs/api_conventions.md index 4ae5146..78ae92e 100644 --- a/docs/api_conventions.md +++ b/docs/api_conventions.md @@ -102,12 +102,16 @@ std::vector // owning; the byte-bool codec returns this const uint8_t* // + length; the byte-bool codec reads this ``` -Each byte represents one bit and must be normalized to `0` or `1`. +Each input byte represents one bit: zero is false and any nonzero value +is true. APIs that produce byte-bools normalize their output to `0` or +`1`. The bit packing, packed-bool I/O, PRG, and GF bit-vector helpers +accept `uint8_t*` directly alongside their existing `bool*` overloads. Do not use `std::vector` in emp-tool library/protocol code. It is bit-packed, has proxy references, has no real `bool*`, and forces -hidden copies. Do not reinterpret byte-bool storage as `bool*`; convert -explicitly if an API requires real `bool` storage. +hidden copies. Do not reinterpret byte-bool storage as `bool*`; pass it +to a byte-bool overload, or convert explicitly if an API requires real +`bool` storage. ## Length and count parameters diff --git a/emp-tool/runtime/core/utils.h b/emp-tool/runtime/core/utils.h index b68d10f..613da03 100644 --- a/emp-tool/runtime/core/utils.h +++ b/emp-tool/runtime/core/utils.h @@ -6,6 +6,7 @@ #include "emp-tool/runtime/core/simd_tier.h" #include #include //https://gcc.gnu.org/gcc-4.9/porting_to.html +#include #include // std::_Exit (fatal abort without running destructors) #include #include "emp-tool/runtime/core/constants.h" @@ -37,17 +38,23 @@ static inline uint32_t bytes_to_bits32(const void* in); // Expand 32 bits into 32 bytes (each 0 or 1). Inverse of bytes_to_bits32. static inline void bits32_to_bytes(uint32_t bits, void* out); -// Pack `len` bools (each byte 0/1) into the first `len` bits of `out`, +// Pack `len` bools into the first `len` bits of `out`, // LSB-first within each byte. Tail-preserve: bits beyond position (len-1) // in the last destination byte are preserved unmodified — callers that // pack into a stack buffer must zero the trailing partial byte before the // call if they care about its contents (e.g. wire-format determinism). // out must hold at least ⌈len/8⌉ bytes. inline void bools_to_bits(void* out, const bool* bools, int64_t len); +template +requires std::is_same_v +inline void bools_to_bits(void* out, const T* bools, int64_t len); // Unpack the first `len` bits of `in` (LSB-first within each byte) into -// `len` bools (each 0/1, byte-sized). bools must hold at least `len` bytes. +// `len` bools (each 0/1). bools must hold at least `len` elements. inline void bits_to_bools(bool* bools, const void* in, int64_t len); +template +requires std::is_same_v +inline void bits_to_bools(T* bools, const void* in, int64_t len); // Value-returning conveniences for fixed-size targets. template diff --git a/emp-tool/runtime/core/utils.hpp b/emp-tool/runtime/core/utils.hpp index 7551913..ce0ae28 100644 --- a/emp-tool/runtime/core/utils.hpp +++ b/emp-tool/runtime/core/utils.hpp @@ -100,7 +100,10 @@ static inline void bits32_to_bytes(uint32_t bits, void *out) { #endif } -inline void bools_to_bits(void *out_, const bool *bools, int64_t len) { +namespace detail { + +template +inline void bools_to_bits_impl(void *out_, const T *bools, int64_t len) { expecting(len >= 0, "bools_to_bits: negative bit count"); uint8_t *out = static_cast(out_); int64_t full32 = len / 32; @@ -110,11 +113,13 @@ inline void bools_to_bits(void *out_, const bool *bools, int64_t len) { } for (int64_t i = full32 * 32; i < len; ++i) { uint8_t mask = (uint8_t)1 << (i % 8); - out[i / 8] = (uint8_t)((out[i / 8] & ~mask) | (((uint8_t)bools[i]) << (i % 8))); + uint8_t bit = static_cast(bools[i] != 0); + out[i / 8] = (uint8_t)((out[i / 8] & ~mask) | (bit << (i % 8))); } } -inline void bits_to_bools(bool *bools, const void *in_, int64_t len) { +template +inline void bits_to_bools_impl(T *bools, const void *in_, int64_t len) { expecting(len >= 0, "bits_to_bools: negative bit count"); const uint8_t *in = static_cast(in_); int64_t full32 = len / 32; @@ -128,6 +133,28 @@ inline void bits_to_bools(bool *bools, const void *in_, int64_t len) { } } +} // namespace detail + +inline void bools_to_bits(void *out, const bool *bools, int64_t len) { + detail::bools_to_bits_impl(out, bools, len); +} + +template +requires std::is_same_v +inline void bools_to_bits(void *out, const T *bools, int64_t len) { + detail::bools_to_bits_impl(out, bools, len); +} + +inline void bits_to_bools(bool *bools, const void *in, int64_t len) { + detail::bits_to_bools_impl(bools, in, len); +} + +template +requires std::is_same_v +inline void bits_to_bools(T *bools, const void *in, int64_t len) { + detail::bits_to_bools_impl(bools, in, len); +} + template inline T bool_to_int(const bool *data) { static_assert(std::is_integral::value, diff --git a/emp-tool/runtime/crypto/f2k.h b/emp-tool/runtime/crypto/f2k.h index 7381aed..8ddff59 100644 --- a/emp-tool/runtime/crypto/f2k.h +++ b/emp-tool/runtime/crypto/f2k.h @@ -2,6 +2,8 @@ #define EMP_F2K_H__ #include "emp-tool/runtime/core/block.h" +#include +#include namespace emp { @@ -37,6 +39,12 @@ inline void vector_inn_prdt_sum_red(block* res, const block* a, const block* b); inline void vector_inn_prdt_sum_red(block* res, const block* a, const bool* b, int64_t sz); template inline void vector_inn_prdt_sum_red(block* res, const block* a, const bool* b); +template +requires std::is_same_v +inline void vector_inn_prdt_sum_red(block* res, const block* a, const T* b, int64_t sz); +template +requires std::is_same_v +inline void vector_inn_prdt_sum_red(block* res, const block* a, const T* b); // Coefficients of the almost-universal hash {seed, seed^2, seed^3, ...}. inline void uni_hash_coeff_gen(block* coeff, block seed, int64_t sz); @@ -55,6 +63,9 @@ class GaloisFieldPacking { public: void packing(block* res, const block* data); void packing(block* res, const bool* data); + template + requires std::is_same_v + void packing(block* res, const T* data); }; } // namespace emp diff --git a/emp-tool/runtime/crypto/f2k.hpp b/emp-tool/runtime/crypto/f2k.hpp index ee59146..d6cc636 100644 --- a/emp-tool/runtime/crypto/f2k.hpp +++ b/emp-tool/runtime/crypto/f2k.hpp @@ -208,25 +208,49 @@ inline void vector_inn_prdt_sum_red(block *res, block const *a, const block *b) vector_inn_prdt_sum_red(res, a, b, N); } -inline void vector_inn_prdt_sum_red(block *res, const block *a, const bool *b, int64_t sz) { +namespace detail { + +template +inline void vector_inn_prdt_sum_red_bits(block *res, const block *a, const T *b, + int64_t sz) { block r0 = zero_block, r1 = zero_block, r2 = zero_block, r3 = zero_block; int64_t i = 0; for (; i + 4 <= sz; i += 4) { - r0 = r0 ^ (a[i ] & select_mask[b[i ]]); - r1 = r1 ^ (a[i+1] & select_mask[b[i+1]]); - r2 = r2 ^ (a[i+2] & select_mask[b[i+2]]); - r3 = r3 ^ (a[i+3] & select_mask[b[i+3]]); + r0 = r0 ^ (a[i ] & select_mask[b[i ] != 0]); + r1 = r1 ^ (a[i+1] & select_mask[b[i+1] != 0]); + r2 = r2 ^ (a[i+2] & select_mask[b[i+2] != 0]); + r3 = r3 ^ (a[i+3] & select_mask[b[i+3] != 0]); } for (; i < sz; ++i) - r0 = r0 ^ (a[i] & select_mask[b[i]]); + r0 = r0 ^ (a[i] & select_mask[b[i] != 0]); *res = (r0 ^ r1) ^ (r2 ^ r3); } +} // namespace detail + +inline void vector_inn_prdt_sum_red(block *res, const block *a, const bool *b, + int64_t sz) { + detail::vector_inn_prdt_sum_red_bits(res, a, b, sz); +} + template inline void vector_inn_prdt_sum_red(block *res, const block *a, const bool *b) { vector_inn_prdt_sum_red(res, a, b, N); } +template +requires std::is_same_v +inline void vector_inn_prdt_sum_red(block *res, const block *a, const T *b, + int64_t sz) { + detail::vector_inn_prdt_sum_red_bits(res, a, b, sz); +} + +template +requires std::is_same_v +inline void vector_inn_prdt_sum_red(block *res, const block *a, const T *b) { + vector_inn_prdt_sum_red(res, a, b, N); +} + inline void uni_hash_coeff_gen(block* coeff, block seed, int64_t sz) { expecting(sz > 0, "uni_hash_coeff_gen: size must be positive"); coeff[0] = seed; @@ -436,6 +460,12 @@ inline void GaloisFieldPacking::packing(block *res, const bool *data) { bools_to_bits(res, data, 128); } +template +requires std::is_same_v +inline void GaloisFieldPacking::packing(block *res, const T *data) { + bools_to_bits(res, data, 128); +} + inline void vector_self_xor(block *sum, block *data, int64_t sz) { block res[4]; res[0] = zero_block; diff --git a/emp-tool/runtime/crypto/prg.h b/emp-tool/runtime/crypto/prg.h index 7ffc391..ff3adcd 100644 --- a/emp-tool/runtime/crypto/prg.h +++ b/emp-tool/runtime/crypto/prg.h @@ -87,34 +87,13 @@ class PRG { public: // consuming a full byte for one bit, an 8x cut in AES work. Inner unpack // uses bits32_to_bytes (SIMD) to expand 4 bytes → 32 bools per call. void random_bool(bool * data, int64_t length) { - expecting(length >= 0, "PRG::random_bool: negative bit count"); - if (length == 0) return; - constexpr int CHUNK_B = 16; // 16 blocks = 2048 bits per pass - block buf[CHUNK_B]; - int64_t produced = 0; - while (produced < length) { - int64_t remaining = length - produced; - int64_t bits_pass = remaining < CHUNK_B * 128 ? remaining : CHUNK_B * 128; - int64_t blocks_pass = (bits_pass + 127) / 128; - random_block(buf, blocks_pass); - const uint8_t *bytes = reinterpret_cast(buf); - int64_t full32 = bits_pass / 32; - for (int64_t i = 0; i < full32; ++i) { - uint32_t b32; - memcpy(&b32, bytes + i * 4, 4); - bits32_to_bytes(b32, data + produced + i * 32); - } - produced += full32 * 32; - int64_t tail_bits = bits_pass - full32 * 32; - if (tail_bits > 0) { - uint32_t b32 = 0; - memcpy(&b32, bytes + full32 * 4, (tail_bits + 7) / 8); - bool tmp[32]; - bits32_to_bytes(b32, tmp); - memcpy(data + produced, tmp, tail_bits); - produced += tail_bits; - } - } + random_bool_impl_(data, length); + } + + template + requires std::is_same_v + void random_bool(T *data, int64_t length) { + random_bool_impl_(data, length); } void random_data_unaligned(void *data, int64_t nbytes) { @@ -215,6 +194,37 @@ class PRG { public: } private: + template + void random_bool_impl_(T *data, int64_t length) { + expecting(length >= 0, "PRG::random_bool: negative bit count"); + if (length == 0) return; + constexpr int CHUNK_B = 16; + block buf[CHUNK_B]; + int64_t produced = 0; + while (produced < length) { + int64_t remaining = length - produced; + int64_t bits_pass = remaining < CHUNK_B * 128 ? remaining : CHUNK_B * 128; + int64_t blocks_pass = (bits_pass + 127) / 128; + random_block(buf, blocks_pass); + const uint8_t *bytes = reinterpret_cast(buf); + int64_t full32 = bits_pass / 32; + for (int64_t i = 0; i < full32; ++i) { + uint32_t b32; + memcpy(&b32, bytes + i * 4, 4); + bits32_to_bytes(b32, data + produced + i * 32); + } + produced += full32 * 32; + int64_t tail_bits = bits_pass - full32 * 32; + if (tail_bits > 0) { + uint32_t b32 = 0; + memcpy(&b32, bytes + full32 * 4, (tail_bits + 7) / 8); + for (int64_t i = 0; i < tail_bits; ++i) + data[produced + i] = static_cast((b32 >> i) & 1); + produced += tail_bits; + } + } + } + uint64_t counter = 0; AES_KEY aes; block key; diff --git a/emp-tool/runtime/io/io_channel.h b/emp-tool/runtime/io/io_channel.h index e010cb9..cbf1809 100644 --- a/emp-tool/runtime/io/io_channel.h +++ b/emp-tool/runtime/io/io_channel.h @@ -249,6 +249,28 @@ class IOChannel { // written before each pack. Whole-byte bytes get fully overwritten by // the SIMD/memcpy path inside bools_to_bits, so they don't need a clear. void send_bool(const bool *data, int64_t length) { + send_bool_impl_(data, length); + } + + template + requires std::is_same_v + void send_bool(const T *data, int64_t length) { + send_bool_impl_(data, length); + } + + void recv_bool(bool *data, int64_t length) { + recv_bool_impl_(data, length); + } + + template + requires std::is_same_v + void recv_bool(T *data, int64_t length) { + recv_bool_impl_(data, length); + } + +private: + template + void send_bool_impl_(const T *data, int64_t length) { expecting(length >= 0, "IOChannel::send_bool: negative bit count"); if (length == 0) return; uint8_t buf[IO_BOOL_CHUNK_SIZE / 8]; @@ -263,7 +285,8 @@ class IOChannel { } } - void recv_bool(bool *data, int64_t length) { + template + void recv_bool_impl_(T *data, int64_t length) { expecting(length >= 0, "IOChannel::recv_bool: negative bit count"); if (length == 0) return; uint8_t buf[IO_BOOL_CHUNK_SIZE / 8]; @@ -277,7 +300,6 @@ class IOChannel { } } -private: // Last traffic direction, for the `rounds` counter. NONE until the first // send/recv so the opening transfer in either direction opens round 1. enum class Dir { NONE, SEND, RECV }; diff --git a/test/runtime/test_block.cpp b/test/runtime/test_block.cpp index 99ea12d..9c21feb 100644 --- a/test/runtime/test_block.cpp +++ b/test/runtime/test_block.cpp @@ -287,26 +287,38 @@ static bool check_bits_bytes_roundtrip() { } static bool check_bools_bits_roundtrip() { - PRG prg; + uint8_t *no_bytes = nullptr; + bools_to_bits(nullptr, nullptr, 0); + bits_to_bools(nullptr, nullptr, 0); + bools_to_bits(nullptr, no_bytes, 0); + bits_to_bools(no_bytes, nullptr, 0); + bool bools_in[4097], bools_out[4097]; for (int len : {1, 7, 8, 9, 31, 32, 33, 127, 128, 1023, 1024, 4097}) { - vector bools_in(len); // byte-bools (0/1) — never vector - for (int i = 0; i < len; ++i) bools_in[i] = (i * 2654435761u) & 1; + vector bytes_in(len); + for (int i = 0; i < len; ++i) { + bytes_in[i] = ((i * 2654435761u) & 1) ? 0xA5 : 0; + bools_in[i] = bytes_in[i] != 0; + } - vector packed((len + 7) / 8, 0xAA); // sentinel byte - bools_to_bits(packed.data(), reinterpret_cast(bools_in.data()), len); + vector packed_bytes((len + 7) / 8, 0xAA); + vector packed_bools((len + 7) / 8, 0xAA); + bools_to_bits(packed_bytes.data(), bytes_in.data(), len); + bools_to_bits(packed_bools.data(), bools_in, len); + if (packed_bytes != packed_bools) return false; - vector out_bytes(len); - bits_to_bools(reinterpret_cast(out_bytes.data()), packed.data(), len); + vector bytes_out(len); + bits_to_bools(bytes_out.data(), packed_bytes.data(), len); + bits_to_bools(bools_out, packed_bytes.data(), len); for (int i = 0; i < len; ++i) { - bool got = out_bytes[i] != 0; - if (got != (bools_in[i] != 0)) return false; + if (bytes_out[i] > 1 || (bytes_out[i] != 0) != bools_in[i]) return false; + if (bools_out[i] != bools_in[i]) return false; } // Tail-byte preservation: bits beyond `len` in the last byte must remain 0xAA. if (len % 8 != 0) { uint8_t mask_below = (uint8_t)((1u << (len % 8)) - 1); uint8_t expected_tail_bits_above = 0xAA & (uint8_t)~mask_below; - uint8_t got_tail_above = packed.back() & (uint8_t)~mask_below; + uint8_t got_tail_above = packed_bytes.back() & (uint8_t)~mask_below; if (got_tail_above != expected_tail_bits_above) return false; } } @@ -326,7 +338,7 @@ static bool run_correctness() { {"sse_trans round-trip", check_sse_trans_roundtrip}, {"sse_trans_n128 parity", check_sse_trans_n128_parity}, {"bytes<->bits32 round-trip", check_bits_bytes_roundtrip}, - {"bools<->bits round-trip", check_bools_bits_roundtrip}, + {"bool/byte-bools<->bits parity", check_bools_bits_roundtrip}, }; bool all = true; for (auto &c : cases) { diff --git a/test/runtime/test_f2k.cpp b/test/runtime/test_f2k.cpp index 99994d2..9ba051e 100644 --- a/test/runtime/test_f2k.cpp +++ b/test/runtime/test_f2k.cpp @@ -20,6 +20,7 @@ #include #include #include +#include #include #include #include @@ -241,35 +242,39 @@ static bool check_vector_inn_prdt_bool(int sz) { PRG prg; vector xs(sz); prg.random_block(xs.data(), sz); - vector bs_bytes(sz); - prg.random_bool(reinterpret_cast(bs_bytes.data()), sz); - const bool *bs = reinterpret_cast(bs_bytes.data()); + vector bs(sz); + prg.random_bool(bs.data(), sz); + auto bools = make_unique(sz); + for (int i = 0; i < sz; ++i) bools[i] = bs[i] != 0; - block got; - vector_inn_prdt_sum_red(&got, xs.data(), bs, sz); + block got_bytes, got_bools; + vector_inn_prdt_sum_red(&got_bytes, xs.data(), bs.data(), sz); + vector_inn_prdt_sum_red(&got_bools, xs.data(), bools.get(), sz); // Reference: XOR of xs[i] for indices where bs[i] = 1. block want = makeBlock(0, 0); for (int i = 0; i < sz; ++i) if (bs[i]) want = want ^ xs[i]; - return blocks_eq(got, want); + return blocks_eq(got_bytes, want) && blocks_eq(got_bools, want); } static bool check_packing_bool() { PRG prg; GaloisFieldPacking pkr; for (int t = 0; t < 16; ++t) { - uint8_t bits_bytes[128]; - prg.random_bool(reinterpret_cast(bits_bytes), 128); - const bool *bits = reinterpret_cast(bits_bytes); - block got; - pkr.packing(&got, bits); + uint8_t bits[128]; + bool bools[128]; + prg.random_bool(bits, 128); + for (int i = 0; i < 128; ++i) bools[i] = bits[i] != 0; + block got_bytes, got_bools; + pkr.packing(&got_bytes, bits); + pkr.packing(&got_bools, bools); // Reference: same identity as block-version, with each X^i contributing // only when bits[i]=1. block want = makeBlock(0, 0); for (int i = 0; i < 128; ++i) if (bits[i]) want = want ^ set_bit(makeBlock(0, 0), i); - if (!blocks_eq(got, want)) return false; + if (!blocks_eq(got_bytes, want) || !blocks_eq(got_bools, want)) return false; } return true; } @@ -293,7 +298,7 @@ static bool run_correctness() { cout << " vector_inn_prdt_sum_* " << (b ? "OK" : "FAIL") << "\n"; bool b2 = true; for (int sz : {1, 7, 64, 1024}) b2 &= check_vector_inn_prdt_bool(sz); - cout << " vector_inn_prdt_sum_red(bool) " << (b2 ? "OK" : "FAIL") << "\n"; + cout << " vector_inn_prdt_sum_red(bits) " << (b2 ? "OK" : "FAIL") << "\n"; bool c = true; for (int sz : {1, 4, 16, 1024}) c &= check_uni_hash_coeff_gen(sz); cout << " uni_hash_coeff_gen " << (c ? "OK" : "FAIL") << "\n"; @@ -303,7 +308,7 @@ static bool run_correctness() { bool d = check_packing(); cout << " GaloisFieldPacking::packing " << (d ? "OK" : "FAIL") << "\n"; bool d2 = check_packing_bool(); - cout << " GaloisFieldPacking::packing(bool) " << (d2 ? "OK" : "FAIL") << "\n"; + cout << " GaloisFieldPacking::packing(bits) " << (d2 ? "OK" : "FAIL") << "\n"; bool e = true; for (int sz : {1, 4, 17, 1024}) e &= check_vector_self_xor(sz); cout << " vector_self_xor " << (e ? "OK" : "FAIL") << "\n"; diff --git a/test/runtime/test_netio.cpp b/test/runtime/test_netio.cpp index 23f4648..3f2b319 100644 --- a/test/runtime/test_netio.cpp +++ b/test/runtime/test_netio.cpp @@ -14,7 +14,9 @@ // regression checks can be reused by IO implementations. #include +#include #include +#include #include "emp-tool/emp-tool.h" @@ -28,6 +30,15 @@ using namespace emp; // ------------------------------------------------------------------------- template static void run_correctness(IO *io, int party, const char *tag) { + uint64_t sent_before = io->send_counter, recv_before = io->recv_counter; + uint8_t *no_bytes = nullptr; + io->send_bool(nullptr, 0); + io->recv_bool(nullptr, 0); + io->send_bool(no_bytes, 0); + io->recv_bool(no_bytes, 0); + expecting(io->send_counter == sent_before && io->recv_counter == recv_before, + "NetIO test: zero-length bool transfer changed counters"); + // Stream of unaligned-byte sends: sends `length` bytes 1000 times in // each direction, with `length` chosen to straddle the 32 KiB sender // staging buffer (NETWORK_STAGING_BUFFER_SIZE/5 + 100) so most send_data calls @@ -62,21 +73,35 @@ static void run_correctness(IO *io, int party, const char *tag) { // Bool packing: 1 MiB of bools sent both aligned and at offset +7 (so // the implementation cannot lean on uint64_t-aligned input). { + constexpr int N = 1024 * 1024; PRG prg(&zero_block); - bool *data = new bool[1024 * 1024]; - bool *data2 = new bool[1024 * 1024]; - prg.random_bool(data, 1024 * 1024); + bool *data = new bool[N]; + bool *data2 = new bool[N]; + vector bytes(N), bytes2(N); + prg.random_bool(data, N); + for (int i = 0; i < N; ++i) bytes[i] = static_cast(data[i]); if (party == ALICE) { - io->send_bool(data, 1024 * 1024); - io->send_bool(data + 7, 1024 * 1024 - 7); + io->send_bool(data, N); + io->send_bool(data + 7, N - 7); + io->send_bool(bytes.data() + 3, N - 3); + io->send_bool(data + 5, N - 5); } else { - io->recv_bool(data2, 1024 * 1024); - expecting(memcmp(data2, data, 1024 * 1024) == 0, + io->recv_bool(data2, N); + expecting(memcmp(data2, data, N) == 0, "NetIO test: aligned bool round-trip mismatch"); - memset(data2, 0, 1024 * 1024); - io->recv_bool(data2 + 7, 1024 * 1024 - 7); - expecting(memcmp(data2 + 7, data + 7, 1024 * 1024 - 7) == 0, + memset(data2, 0, N); + io->recv_bool(data2 + 7, N - 7); + expecting(memcmp(data2 + 7, data + 7, N - 7) == 0, "NetIO test: unaligned bool round-trip mismatch"); + memset(data2, 0, N); + io->recv_bool(data2 + 3, N - 3); + for (int i = 3; i < N; ++i) + expecting(data2[i] == data[i], + "NetIO test: byte-bool send mismatch"); + io->recv_bool(bytes2.data() + 5, N - 5); + for (int i = 5; i < N; ++i) + expecting((bytes2[i] != 0) == data[i], + "NetIO test: byte-bool receive mismatch"); } delete[] data; delete[] data2; diff --git a/test/runtime/test_prg.cpp b/test/runtime/test_prg.cpp index cfa7bcc..b183b54 100644 --- a/test/runtime/test_prg.cpp +++ b/test/runtime/test_prg.cpp @@ -214,20 +214,43 @@ static bool check_random_bool_is_0_or_1() { PRG p; for (int len : {1, 7, 32, 128, 1023, 4096}) { vector buf(len); - // random_bool writes via bool*; 1 byte per bool, value 0 or 1. - p.random_bool(reinterpret_cast(buf.data()), len); + p.random_bool(buf.data(), len); for (int i = 0; i < len; ++i) if (buf[i] > 1) return false; } return true; } +static bool check_random_bool_representations_match() { + block seed = makeBlock(0x1234, 0x5678); + PRG empty(&seed); + uint8_t *no_bytes = nullptr; + empty.random_bool(nullptr, 0); + empty.random_bool(no_bytes, 0); + if (empty.position() != 0) return false; + bool bools[4097]; + for (int len : {0, 1, 7, 31, 32, 33, 127, 128, 129, 2048, 2049, 4097}) { + PRG a(&seed), b(&seed); + vector bytes(len); + a.random_bool(bools, len); + b.random_bool(bytes.data(), len); + for (int i = 0; i < len; ++i) + if (bytes[i] > 1 || (bytes[i] != 0) != bools[i]) return false; + if (a.position() != b.position()) return false; + block next_a, next_b; + a.random_block(&next_a, 1); + b.random_block(&next_b, 1); + if (!blocks_eq(next_a, next_b)) return false; + } + return true; +} + static bool check_random_bool_distribution() { // Weak distribution check: mean bit ≈ 0.5 over a large sample. PRG p; const int N = 1 << 18; vector buf(N); - p.random_bool(reinterpret_cast(buf.data()), N); + p.random_bool(buf.data(), N); int64_t ones = 0; for (int i = 0; i < N; ++i) ones += buf[i]; double mean = (double)ones / N; @@ -276,6 +299,7 @@ static bool run_correctness() { {"random_data_unaligned vs aligned ref", check_random_data_unaligned}, {"random_data_unaligned counter", check_random_data_unaligned_counter}, {"random_bool ∈ {0,1}", check_random_bool_is_0_or_1}, + {"random_bool bool/byte parity", check_random_bool_representations_match}, {"random_bool mean ~ 0.5", check_random_bool_distribution}, {"UniformRandomBitGenerator interface", check_uniform_engine}, {"negative lengths rejected", check_negative_lengths_rejected}, From 8322e4a6e4732d3ed7b3988db25e99e43483092f Mon Sep 17 00:00:00 2001 From: Xiao Wang Date: Sat, 22 Aug 2026 08:38:33 -0500 Subject: [PATCH 2/7] Add pre-connect TCP socket buffer options Expose explicit send and receive socket-buffer sizes through NetIO and TLSConfig. Apply them before listen/connect and the TLS handshake, propagate them to NetIO siblings, and leave operating-system defaults unchanged when both values are zero. Provide an opt-in bandwidth/RTT helper that rounds the bandwidth-delay product to a power-of-two tier and reject explicit sizes the kernel cannot provide. --- docs/io_channel.md | 47 +++++++-- emp-tool/runtime/io/net_io_channel.h | 33 ++++-- emp-tool/runtime/io/tcp_socket.h | 144 +++++++++++++++++++++++++-- emp-tool/runtime/io/tls_io_channel.h | 19 ++-- test/runtime/test_netio.cpp | 110 +++++++++++++++++++- test/runtime/test_tlsio.cpp | 51 ++++++++++ 6 files changed, 373 insertions(+), 31 deletions(-) diff --git a/docs/io_channel.md b/docs/io_channel.md index 069e4b3..180e9ca 100644 --- a/docs/io_channel.md +++ b/docs/io_channel.md @@ -87,17 +87,50 @@ are per-direction snapshots for diagnostics. All three assert that `make_sibling()`, calling it **serially and in the same order on both parties** (its accept/connect pairing is FIFO on the shared listener — concurrent `make_sibling()` from multiple threads is not deterministic). - Do *not* rely on closing every channel and reopening a new one on the - same port as the coordination mechanism. That reopen path is supported - and race-free — each connection is only considered established once the - peer has actually `accept()`ed it (a one-byte accept acknowledgement in - `tcp_socket.h` protects against a `connect()` landing on the previous, - now-stale listener) — but the anchor + `make_sibling` pattern is simpler - and avoids the reconnect entirely. + Closing all channels and reopening the same port is supported, but sibling + channels avoid reconnecting. - **`TraceIO`** (`trace_io.h`): an `IOChannel` that tees every wire byte to `.send` / `.recv` files for diff-based wire-equivalence checks; see `test_mode.md`. +## TCP socket buffers + +`tcp::SocketOptions` sets the send and receive buffer sizes before +`listen()` or `connect()`: + +```cpp +tcp::SocketOptions options; +options.send_buffer_size = 16 * 1024 * 1024; +options.receive_buffer_size = 16 * 1024 * 1024; + +auto io = NetIO::connect(peer, port, options); +``` + +When the path capacity and round-trip time are known, the helper sizes both +directions to the first power-of-two tier at or above the bandwidth-delay +product, with a 256 KiB minimum: + +```cpp +using namespace std::chrono_literals; + +auto options = tcp::SocketOptions::for_bandwidth_and_rtt( + 400'000'000, 100ms); // 8 MiB +auto io = NetIO::connect(peer, port, options); +``` + +Use the direct byte fields when the two directions need different sizes. + +Pass the same options to `NetIO::listen`; `make_sibling()` propagates them to +each new connection. `TLSIO` takes them through `TLSConfig::socket_options`. +A listening endpoint applies the options before `listen()` and verifies them +on each accepted socket. TLSIO applies the options before the TLS handshake. +The adopted-socket TLS constructor requires default socket options. +A zero size keeps the operating-system default. An explicit request fails if +the kernel caps the buffer below the requested size. Larger requests may require +raising `net.core.wmem_max` / `net.core.rmem_max` on Linux or +`kern.ipc.maxsockbuf` on macOS. On Linux, setting `SO_RCVBUF` disables TCP +receive-buffer autotuning for that socket. + ## TLS variant `TLSIO` (in `emp-tool/runtime/io/tls_io_channel.h`) is another `IOChannel` diff --git a/emp-tool/runtime/io/net_io_channel.h b/emp-tool/runtime/io/net_io_channel.h index 36e9318..96af7e3 100644 --- a/emp-tool/runtime/io/net_io_channel.h +++ b/emp-tool/runtime/io/net_io_channel.h @@ -41,6 +41,7 @@ class NetIO : public IOChannel { public: // Endpoint info retained so a duplex sibling can be spawned (make_sibling). std::string addr_; // peer address (empty when this is a server) int port_ = -1; + tcp::SocketOptions socket_options_; // Send-side state (stdio "wb" stream + app-level coalescing buffer). FILE *stream = nullptr; @@ -59,7 +60,12 @@ class NetIO : public IOChannel { public: // time per channel; threaded consumers take a sibling channel each // (make_sibling). Races are not detected at runtime — use TSan. - NetIO(const char *address, int port, bool quiet = false) : quiet(quiet) { + NetIO(const char *address, int port, bool quiet = false) + : NetIO(address, port, tcp::SocketOptions{}, quiet) {} + + NetIO(const char *address, int port, + const tcp::SocketOptions &socket_options, bool quiet = false) + : quiet(quiet), socket_options_(socket_options) { expecting(port >= 0 && port <= 65535, "NetIO: invalid port number"); @@ -67,10 +73,13 @@ class NetIO : public IOChannel { public: addr_ = address ? address : ""; port_ = port; if (is_server) { - listener = std::make_shared(tcp::open_listener(port)); - init_from_sock(tcp::accept_one_confirmed(listener->fd)); + listener = std::make_shared( + tcp::open_listener(port, socket_options_)); + init_from_sock(tcp::accept_one_confirmed(listener->fd, + socket_options_)); } else { - init_from_sock(tcp::client_connect_confirmed(address, port)); + init_from_sock(tcp::client_connect_confirmed(address, port, + socket_options_)); } if (!quiet) std::cout << "connected\n"; } @@ -82,9 +91,19 @@ class NetIO : public IOChannel { public: static std::unique_ptr listen(int port, bool quiet = false) { return std::make_unique(nullptr, port, quiet); } + static std::unique_ptr listen(int port, + const tcp::SocketOptions &socket_options, + bool quiet = false) { + return std::make_unique(nullptr, port, socket_options, quiet); + } static std::unique_ptr connect(const char *address, int port, bool quiet = false) { return std::make_unique(address, port, quiet); } + static std::unique_ptr connect( + const char *address, int port, const tcp::SocketOptions &socket_options, + bool quiet = false) { + return std::make_unique(address, port, socket_options, quiet); + } // Open another channel to the same peer and port. Related server channels // share the listener, so make_sibling() may be called repeatedly on the @@ -92,15 +111,17 @@ class NetIO : public IOChannel { public: // port. The listener closes when the last related server NetIO is destroyed. std::unique_ptr make_sibling() const { if (!is_server) - return connect(addr_.c_str(), port_, /*quiet=*/true); + return connect(addr_.c_str(), port_, socket_options_, /*quiet=*/true); expecting(listener != nullptr, "NetIO::make_sibling requires a server listener"); - int sibling_sock = tcp::accept_one_confirmed(listener->fd); + int sibling_sock = tcp::accept_one_confirmed(listener->fd, + socket_options_); auto sibling = std::make_unique(sibling_sock, /*quiet=*/true); sibling->is_server = true; // preserve sync()'s server/client ordering sibling->port_ = port_; + sibling->socket_options_ = socket_options_; sibling->listener = listener; return sibling; } diff --git a/emp-tool/runtime/io/tcp_socket.h b/emp-tool/runtime/io/tcp_socket.h index 520915c..008e90a 100644 --- a/emp-tool/runtime/io/tcp_socket.h +++ b/emp-tool/runtime/io/tcp_socket.h @@ -10,9 +10,12 @@ #include #include +#include #include #include +#include #include +#include #include #include "emp-tool/runtime/core/error.h" @@ -20,9 +23,96 @@ namespace emp { namespace tcp { +struct SocketOptions { + int send_buffer_size = 0; + int receive_buffer_size = 0; + + static SocketOptions for_bandwidth_and_rtt( + std::uint64_t bandwidth_bits_per_second, + std::chrono::microseconds round_trip_time) { + expecting(bandwidth_bits_per_second > 0, + "tcp: bandwidth must be positive"); + expecting(round_trip_time.count() > 0, + "tcp: round-trip time must be positive"); + + constexpr std::uint64_t byte_scale = 8'000'000; + constexpr std::uint64_t minimum_buffer = 256 * 1024; + constexpr std::uint64_t maximum_buffer = + static_cast(std::numeric_limits::max()); + const auto rtt_microseconds = + static_cast(round_trip_time.count()); + const auto maximum_product = maximum_buffer * byte_scale; + expecting(bandwidth_bits_per_second <= + maximum_product / rtt_microseconds, + "tcp: bandwidth-delay product exceeds socket buffer range"); + + const auto product = bandwidth_bits_per_second * rtt_microseconds; + const auto bdp_bytes = (product + byte_scale - 1) / byte_scale; + auto buffer_size = minimum_buffer; + while (buffer_size < bdp_bytes) { + expecting(buffer_size <= maximum_buffer / 2, + "tcp: rounded bandwidth-delay product exceeds socket buffer range"); + buffer_size *= 2; + } + + const auto size = static_cast(buffer_size); + return SocketOptions{size, size}; + } +}; + +inline void verify_socket_buffer(int sock, int option, int requested, + const char *name) { + expecting(requested >= 0, [&] { + return std::string("tcp: ") + name + " cannot be negative"; + }); + if (requested == 0) return; + + int actual = 0; + socklen_t length = sizeof(actual); + expecting(::getsockopt(sock, SOL_SOCKET, option, &actual, &length) == 0, + [&] { + return std::string("tcp: getsockopt(") + name + "): " + + std::strerror(errno); + }); +#ifdef __linux__ + // Linux reports twice the user-visible socket buffer size. + actual /= 2; +#endif + expecting(actual >= requested, [&] { + return std::string("tcp: ") + name + " requested " + + std::to_string(requested) + " bytes, kernel provided " + + std::to_string(actual); + }); +} + +inline void set_socket_buffer(int sock, int option, int requested, + const char *name) { + expecting(requested >= 0, [&] { + return std::string("tcp: ") + name + " cannot be negative"; + }); + if (requested == 0) return; + + expecting(::setsockopt(sock, SOL_SOCKET, option, &requested, + sizeof(requested)) == 0, [&] { + return std::string("tcp: setsockopt(") + name + "): " + + std::strerror(errno); + }); + verify_socket_buffer(sock, option, requested, name); +} + +inline void apply_socket_options(int sock, const SocketOptions &options) { + set_socket_buffer(sock, SO_SNDBUF, options.send_buffer_size, "SO_SNDBUF"); + set_socket_buffer(sock, SO_RCVBUF, options.receive_buffer_size, "SO_RCVBUF"); +} + +inline void verify_socket_options(int sock, const SocketOptions &options) { + verify_socket_buffer(sock, SO_SNDBUF, options.send_buffer_size, "SO_SNDBUF"); + verify_socket_buffer(sock, SO_RCVBUF, options.receive_buffer_size, "SO_RCVBUF"); +} + // Bind and listen without accepting. Related server NetIO channels share this // descriptor so repeated sibling connections use the same listener. -inline int open_listener(int port) { +inline int open_listener(int port, const SocketOptions &options) { struct sockaddr_in serv; std::memset(&serv, 0, sizeof(serv)); serv.sin_family = AF_INET; @@ -32,6 +122,7 @@ inline int open_listener(int port) { expecting(listener >= 0, [&] { return std::string("tcp: socket: ") + std::strerror(errno); }); + apply_socket_options(listener, options); int reuse = 1; ::setsockopt(listener, SOL_SOCKET, SO_REUSEADDR, (const char *)&reuse, sizeof(reuse)); expecting(::bind(listener, (struct sockaddr *)&serv, @@ -50,6 +141,10 @@ inline int open_listener(int port) { return listener; } +inline int open_listener(int port) { + return open_listener(port, SocketOptions{}); +} + // Shared RAII owner for a listening socket. NetIO siblings share one handle so // any related channel can accept another sibling and the listener closes when // the last related server channel is destroyed. @@ -86,11 +181,17 @@ inline int accept_one(int listener) { return s; } +inline int accept_one(int listener, const SocketOptions &options) { + int s = accept_one(listener); + verify_socket_options(s, options); + return s; +} + // accept_one plus a one-byte application acknowledgement, so the client can // confirm this accept() actually ran (see kAcceptAck). Raw ::send, so it does // not touch IOChannel byte counters or any Fiat-Shamir transcript. -inline int accept_one_confirmed(int listener) { - int s = accept_one(listener); +inline int accept_one_confirmed(int listener, const SocketOptions &options) { + int s = accept_one(listener, options); ssize_t n; do { n = ::send(s, &kAcceptAck, 1, 0); @@ -99,15 +200,23 @@ inline int accept_one_confirmed(int listener) { return s; } +inline int accept_one_confirmed(int listener) { + return accept_one_confirmed(listener, SocketOptions{}); +} + // One-shot compatibility helper used by transports that need only one // connection, such as TLSIO. -inline int server_listen(int port) { - int listener = open_listener(port); - int s = accept_one(listener); +inline int server_listen(int port, const SocketOptions &options) { + int listener = open_listener(port, options); + int s = accept_one(listener, options); ::close(listener); return s; } +inline int server_listen(int port) { + return server_listen(port, SocketOptions{}); +} + // Connect to address:port, retrying on failure with a 1 ms backoff so // the server side has time to come up. Capped at ~60 s of total retry // time to catch a permanently-down peer instead of hanging the caller. @@ -116,7 +225,8 @@ inline int server_listen(int port) { // listener the server no longer accepts on — is treated as a failure and // retried, so "connected" always means the peer's accept() ran. inline int client_connect_impl(const char *address, int port, - bool wait_for_accept) { + bool wait_for_accept, + const SocketOptions &options) { struct sockaddr_in dest; std::memset(&dest, 0, sizeof(dest)); dest.sin_family = AF_INET; @@ -128,6 +238,7 @@ inline int client_connect_impl(const char *address, int port, expecting(s >= 0, [&] { return std::string("tcp: socket: ") + std::strerror(errno); }); + apply_socket_options(s, options); if (::connect(s, (struct sockaddr *)&dest, sizeof(struct sockaddr)) == 0) { if (!wait_for_accept) return s; @@ -149,17 +260,32 @@ inline int client_connect_impl(const char *address, int port, error(msg.c_str()); } +inline int client_connect_impl(const char *address, int port, + bool wait_for_accept) { + return client_connect_impl(address, port, wait_for_accept, SocketOptions{}); +} + // Plain connect: returns as soon as the TCP handshake completes (the peer may // not have accept()ed yet). Used by transports with their own post-connect // handshake, e.g. TLSIO. +inline int client_connect(const char *address, int port, + const SocketOptions &options) { + return client_connect_impl(address, port, /*wait_for_accept=*/false, options); +} + inline int client_connect(const char *address, int port) { - return client_connect_impl(address, port, /*wait_for_accept=*/false); + return client_connect(address, port, SocketOptions{}); } // Connect and wait for the server's accept acknowledgement — safe against a // connection landing on a stale listener (see kAcceptAck). Used by NetIO. +inline int client_connect_confirmed(const char *address, int port, + const SocketOptions &options) { + return client_connect_impl(address, port, /*wait_for_accept=*/true, options); +} + inline int client_connect_confirmed(const char *address, int port) { - return client_connect_impl(address, port, /*wait_for_accept=*/true); + return client_connect_confirmed(address, port, SocketOptions{}); } inline void set_nodelay(int sock) { diff --git a/emp-tool/runtime/io/tls_io_channel.h b/emp-tool/runtime/io/tls_io_channel.h index eff3c94..ea8d9bc 100644 --- a/emp-tool/runtime/io/tls_io_channel.h +++ b/emp-tool/runtime/io/tls_io_channel.h @@ -40,10 +40,7 @@ namespace emp { struct TLSConfig { // Role on the TLS handshake. Independent of the address-vs-port - // "is_server" bit on NetIO: a TCP-acceptor side could in principle - // be a TLS client (reverse-direction handshake), though in practice - // the bits track each other. TLSIO defaults to is_tls_server = - // (address == nullptr) when constructed via the (addr, port) ctor. + // role used to establish the underlying TCP connection. bool is_tls_server = false; // PEM file paths. cert_pem_path + key_pem_path are required when @@ -64,6 +61,8 @@ struct TLSConfig { // "" leaves the library default TLS 1.3 ciphersuite list untouched. std::string ciphersuites; + + tcp::SocketOptions socket_options; }; namespace tls_detail { @@ -124,8 +123,8 @@ class TLSIO : public IOChannel { public: // mutates internal state on every read/write): one thread at a time // per channel. Races are not detected at runtime — use TSan. - // (addr == nullptr) → TCP listener; otherwise TCP client. is_tls_server - // defaults to mirror is_server but TLSConfig overrides if explicitly set. + // (addr == nullptr) → TCP listener; otherwise TCP client. The TLS role + // comes from TLSConfig and is independent of the TCP role. TLSIO(const char *address, int port, const TLSConfig &cfg, bool quiet = false) : quiet(quiet) { expecting(port >= 0 && port <= 65535, @@ -133,8 +132,9 @@ class TLSIO : public IOChannel { public: tls_detail::install_sigpipe_ignore_once(); is_server = (address == nullptr); is_tls_server = cfg.is_tls_server; - init_from_sock(is_server ? tcp::server_listen(port) - : tcp::client_connect(address, port), + init_from_sock(is_server ? tcp::server_listen(port, cfg.socket_options) + : tcp::client_connect(address, port, + cfg.socket_options), cfg); if (!quiet) std::cout << "TLS connected\n"; } @@ -155,6 +155,9 @@ class TLSIO : public IOChannel { public: TLSIO(int existing_sock, bool is_tls_server, const TLSConfig &cfg, bool quiet = true) : quiet(quiet), is_tls_server(is_tls_server) { + expecting(cfg.socket_options.send_buffer_size == 0 && + cfg.socket_options.receive_buffer_size == 0, + "TLSIO: socket buffer options cannot be applied to an already-connected socket"); tls_detail::install_sigpipe_ignore_once(); is_server = false; init_from_sock(existing_sock, cfg); diff --git a/test/runtime/test_netio.cpp b/test/runtime/test_netio.cpp index 3f2b319..f7c56cb 100644 --- a/test/runtime/test_netio.cpp +++ b/test/runtime/test_netio.cpp @@ -9,6 +9,9 @@ // flush() drain outbound only (no peer coupling) // sync() 1-byte ping/pong handshake // make_sibling() more connections on the same port +// tcp::SocketOptions pre-handshake socket buffer sizing +// SocketOptions::for_bandwidth_and_rtt +// derive buffers from a known path // // Test functions below are templated on the IO type so correctness and // regression checks can be reused by IO implementations. @@ -16,13 +19,31 @@ #include #include #include +#include #include +#include +#include + #include "emp-tool/emp-tool.h" using namespace std; using namespace emp; +template +static bool dies(F &&f) { + pid_t pid = fork(); + expecting(pid >= 0, "NetIO test: fork failed"); + if (pid == 0) { + std::freopen("/dev/null", "w", stderr); + f(); + _exit(0); + } + int status = 0; + waitpid(pid, &status, 0); + return !(WIFEXITED(status) && WEXITSTATUS(status) == 0); +} + // ------------------------------------------------------------------------- // run_correctness(): byte stream round-trip at unaligned offsets, then bool // packing round-trip at unaligned bool offsets. Each side asserts on the @@ -209,6 +230,90 @@ static void run_sibling_regression(int port, int party) { if (party == ALICE) cout << "NetIO shared-listener regression: OK\n"; } +static int socket_buffer_size(int sock, int option) { + int size = 0; + socklen_t length = sizeof(size); + expecting(::getsockopt(sock, SOL_SOCKET, option, &size, &length) == 0, + "NetIO test: getsockopt failed"); +#ifdef __linux__ + size /= 2; +#endif + return size; +} + +static void expect_socket_options(const NetIO &io, + const tcp::SocketOptions &options) { + expecting(socket_buffer_size(io.sock, SO_SNDBUF) >= options.send_buffer_size, + "NetIO test: send buffer option was not applied"); + expecting(socket_buffer_size(io.sock, SO_RCVBUF) >= options.receive_buffer_size, + "NetIO test: receive buffer option was not applied"); +} + +static void run_socket_options_factory_regression() { + int (*open_listener_legacy)(int) = tcp::open_listener; + int (*accept_one_confirmed_legacy)(int) = tcp::accept_one_confirmed; + int (*server_listen_legacy)(int) = tcp::server_listen; + int (*client_connect_impl_legacy)(const char *, int, bool) = + tcp::client_connect_impl; + int (*client_connect_legacy)(const char *, int) = tcp::client_connect; + int (*client_connect_confirmed_legacy)(const char *, int) = + tcp::client_connect_confirmed; + (void)open_listener_legacy; + (void)accept_one_confirmed_legacy; + (void)server_listen_legacy; + (void)client_connect_impl_legacy; + (void)client_connect_legacy; + (void)client_connect_confirmed_legacy; + + const auto short_path = tcp::SocketOptions::for_bandwidth_and_rtt( + 400'000'000, std::chrono::microseconds(450)); + expecting(short_path.send_buffer_size == 256 * 1024 && + short_path.receive_buffer_size == 256 * 1024, + "NetIO test: short-path buffer tier mismatch"); + + const auto wan_path = tcp::SocketOptions::for_bandwidth_and_rtt( + 400'000'000, std::chrono::milliseconds(100)); + expecting(wan_path.send_buffer_size == 8 * 1024 * 1024 && + wan_path.receive_buffer_size == 8 * 1024 * 1024, + "NetIO test: WAN buffer tier mismatch"); + + expecting(dies([] { + int sock = ::socket(AF_INET, SOCK_STREAM, 0); + expecting(sock >= 0, "NetIO test: socket failed"); + tcp::SocketOptions unavailable; + unavailable.send_buffer_size = std::numeric_limits::max(); + tcp::verify_socket_options(sock, unavailable); + }), "NetIO test: unavailable socket buffer was not rejected"); +} + +static void run_socket_options_regression(int port, int party) { + tcp::SocketOptions options; + options.send_buffer_size = 128 * 1024; + options.receive_buffer_size = 128 * 1024; + + auto primary = party == ALICE ? NetIO::listen(port, options, true) + : NetIO::connect(peer_ip(), port, options, true); + expect_socket_options(*primary, options); + if (party == ALICE) { + expecting(socket_buffer_size(primary->listener->fd, SO_SNDBUF) >= + options.send_buffer_size, + "NetIO test: listener send buffer option was not applied"); + expecting(socket_buffer_size(primary->listener->fd, SO_RCVBUF) >= + options.receive_buffer_size, + "NetIO test: listener receive buffer option was not applied"); + } + + auto sibling = primary->make_sibling(); + sibling->sync(); + expect_socket_options(*sibling, options); + primary.reset(); + + auto next = sibling->make_sibling(); + next->sync(); + expect_socket_options(*next, options); + if (party == ALICE) cout << "NetIO socket-options regression: OK\n"; +} + // Zero-length send/recv are documented no-ops (docs/api_conventions.md): // no counter, round, or flush-state mutation, no transport call, and a // null pointer is fine at count zero. Purely local — both parties run it @@ -241,8 +346,11 @@ int main(int argc, char **argv) { int port, party; party = parse_party(argv); port = peer_port(); + run_socket_options_factory_regression(); - // Four contiguous ports: main, two send-only cases, sibling regression. + // Five contiguous ports: main, two send-only cases, sibling regression, + // socket-options regression. run_suite(port, party, "NetIO"); run_sibling_regression(port + 3, party); + run_socket_options_regression(port + 4, party); } diff --git a/test/runtime/test_tlsio.cpp b/test/runtime/test_tlsio.cpp index e5c7c98..cb19132 100644 --- a/test/runtime/test_tlsio.cpp +++ b/test/runtime/test_tlsio.cpp @@ -8,6 +8,7 @@ // send_bool / recv_bool packed via bools_to_bits (inherited) // flush() drain outbound coalescing buffer // sync() 1-byte ping/pong handshake +// TLSConfig::socket_options pre-handshake socket buffer sizing // // Same flush contract and thread-safety rules as NetIO. The test // mirrors test_netio.cpp's correctness + send-only regression suite, @@ -39,8 +40,10 @@ #include #include +#include #include #include +#include #include #include @@ -53,6 +56,20 @@ using namespace std; using namespace emp; +template +static bool dies(F &&f) { + pid_t pid = fork(); + expecting(pid >= 0, "TLSIO test: fork failed"); + if (pid == 0) { + std::freopen("/dev/null", "w", stderr); + f(); + _exit(0); + } + int status = 0; + waitpid(pid, &status, 0); + return !(WIFEXITED(status) && WEXITSTATUS(status) == 0); +} + static const char *CA_CERT = "/tmp/emp_tlsio_test_ca_cert.pem"; static const char *ALICE_CERT = "/tmp/emp_tlsio_test_alice_cert.pem"; static const char *ALICE_KEY = "/tmp/emp_tlsio_test_alice_key.pem"; @@ -193,9 +210,41 @@ static TLSConfig make_cfg(int party) { cfg.ca_pem_path = CA_CERT; // both trust the same CA cfg.require_peer_cert = true; // exercise the mTLS path cfg.insecure_skip_verify = false; + cfg.socket_options.send_buffer_size = 128 * 1024; + cfg.socket_options.receive_buffer_size = 128 * 1024; return cfg; } +static int socket_buffer_size(int sock, int option) { + int size = 0; + socklen_t length = sizeof(size); + expecting(::getsockopt(sock, SOL_SOCKET, option, &size, &length) == 0, + "TLSIO test: getsockopt failed"); +#ifdef __linux__ + size /= 2; +#endif + return size; +} + +static void expect_socket_options(const TLSIO &io, + const tcp::SocketOptions &options) { + expecting(socket_buffer_size(io.sock, SO_SNDBUF) >= options.send_buffer_size, + "TLSIO test: send buffer option was not applied"); + expecting(socket_buffer_size(io.sock, SO_RCVBUF) >= options.receive_buffer_size, + "TLSIO test: receive buffer option was not applied"); +} + +static void run_adopted_socket_options_rejection() { + expecting(dies([] { + int sockets[2]; + expecting(::socketpair(AF_UNIX, SOCK_STREAM, 0, sockets) == 0, + "TLSIO test: socketpair failed"); + TLSConfig cfg; + cfg.socket_options.receive_buffer_size = 128 * 1024; + TLSIO io(sockets[0], true, cfg, true); + }), "TLSIO test: adopted socket accepted pre-handshake options"); +} + // ------------------------------------------------------------------------- // run_correctness(): byte stream round-trip at unaligned offsets, then bool // packing round-trip at unaligned bool offsets. Same shape as test_netio. @@ -309,6 +358,7 @@ int main(int argc, char **argv) { int port, party; party = parse_party(argv); port = peer_port(); + run_adopted_socket_options_rejection(); // PKI handoff. ALICE (party 1, the TCP listener) builds the whole // CA + ALICE + BOB PKI in memory and writes 5 PEM files; BOB polls @@ -321,6 +371,7 @@ int main(int argc, char **argv) { const TLSConfig cfg = make_cfg(party); TLSIO *io = new TLSIO(party == ALICE ? nullptr : peer_ip(), port, cfg, true); + expect_socket_options(*io, cfg.socket_options); run_correctness(io, party); run_send_only_regression(port, party); delete io; From fa0fedda407f264dc2152ff1bcaad9b2fb010d82 Mon Sep 17 00:00:00 2001 From: Xiao Wang Date: Sat, 22 Aug 2026 08:38:33 -0500 Subject: [PATCH 3/7] Remove IO sync and report flush failures Remove the transport-level sync hook and preserve connection coverage with explicit protocol round trips. Callers that need a peer barrier must exchange a protocol message; flush only drains local output. Fail on NetIO stdio flush errors. Make TraceIO reject null transports, flush its trace before the wrapped channel, create trace files with mode 0600, and report trace write and flush failures. --- bench/bench_netio.cpp | 1 - bench/bench_tlsio.cpp | 1 - docs/io_channel.md | 3 - docs/test_mode.md | 7 +- emp-tool/runtime/io/io_channel.h | 4 - emp-tool/runtime/io/net_io_channel.h | 19 +--- emp-tool/runtime/io/tls_io_channel.h | 24 ++--- emp-tool/runtime/io/trace_io.h | 53 ++++++++--- test/CMakeLists.txt | 1 + test/runtime/test_netio.cpp | 41 +++++++-- test/runtime/test_tlsio.cpp | 1 - test/runtime/test_traceio.cpp | 127 +++++++++++++++++++++++++++ 12 files changed, 216 insertions(+), 66 deletions(-) create mode 100644 test/runtime/test_traceio.cpp diff --git a/bench/bench_netio.cpp b/bench/bench_netio.cpp index 7be9c56..f0e24bc 100644 --- a/bench/bench_netio.cpp +++ b/bench/bench_netio.cpp @@ -7,7 +7,6 @@ // send_block / recv_block block-typed wrapper // send_bool / recv_bool packed via bools_to_bits // flush() drain outbound only (no peer coupling) -// sync() 1-byte ping/pong handshake // // Benchmark below runs the loopback throughput sweep only. Correctness and // regression coverage lives in test/test_netio.cpp. diff --git a/bench/bench_tlsio.cpp b/bench/bench_tlsio.cpp index c4377b1..544dd43 100644 --- a/bench/bench_tlsio.cpp +++ b/bench/bench_tlsio.cpp @@ -7,7 +7,6 @@ // send_block / recv_block block-typed wrapper (inherited) // send_bool / recv_bool packed via bools_to_bits (inherited) // flush() drain outbound coalescing buffer -// sync() 1-byte ping/pong handshake // // Same flush contract and thread-safety rules as NetIO. This benchmark runs // the TLSIO loopback throughput sweep with the same PKI setup as the test. diff --git a/docs/io_channel.md b/docs/io_channel.md index 180e9ca..d68ef69 100644 --- a/docs/io_channel.md +++ b/docs/io_channel.md @@ -69,9 +69,6 @@ are per-direction snapshots for diagnostics. All three assert that ## Other base surface -- **`sync()`**: optional 1-byte ping/pong handshake to confirm both - directions are alive. NetIO implements it; the base default is a - no-op. - **Telemetry**: the base tracks `send_counter` / `recv_counter` / `rounds` / `flushes_count`; `get_statistics_string()` renders them for logging (`~NetIO` prints it unless constructed `quiet`). diff --git a/docs/test_mode.md b/docs/test_mode.md index 8dd0dad..df6498e 100644 --- a/docs/test_mode.md +++ b/docs/test_mode.md @@ -124,10 +124,9 @@ in production paths. `TraceIO` wraps any `IOChannel*` and writes a copy of every wire byte to two files: `.send` and `.recv`. Bytes are -delivered to the underlying channel either before (recv) or -synchronously (send) with the file write, so a crash mid-write -leaves a trace prefix that still matches what the peer didn't yet -see. +copied before outbound delivery and after inbound delivery. Trace +files are created with mode `0600`; `TraceIO::flush()` flushes both +files before flushing the wrapped channel. ```cpp NetIO* under = new NetIO(...); diff --git a/emp-tool/runtime/io/io_channel.h b/emp-tool/runtime/io/io_channel.h index cbf1809..5572fb4 100644 --- a/emp-tool/runtime/io/io_channel.h +++ b/emp-tool/runtime/io/io_channel.h @@ -87,10 +87,6 @@ class IOChannel { // is a no-op for transports with nothing to flush. virtual void flush() {} - // Optional wire-level handshake (e.g. 1-byte ping/pong). Default - // no-op for transports that don't need one. - virtual void sync() {} - // Turn on Fiat-Shamir transcript hashing. `send_first` selects which // of the two H(_‖H(_)) formulas this side computes, so both parties // produce the same digest value — exactly one party should pass true. diff --git a/emp-tool/runtime/io/net_io_channel.h b/emp-tool/runtime/io/net_io_channel.h index 96af7e3..21962e5 100644 --- a/emp-tool/runtime/io/net_io_channel.h +++ b/emp-tool/runtime/io/net_io_channel.h @@ -119,7 +119,7 @@ class NetIO : public IOChannel { public: socket_options_); auto sibling = std::make_unique(sibling_sock, /*quiet=*/true); - sibling->is_server = true; // preserve sync()'s server/client ordering + sibling->is_server = true; sibling->port_ = port_; sibling->socket_options_ = socket_options_; sibling->listener = listener; @@ -170,19 +170,6 @@ class NetIO : public IOChannel { public: void set_nodelay() { tcp::set_nodelay(sock); } void set_delay() { tcp::set_delay(sock); } - // 1-byte ping/pong handshake to verify both directions are alive. - void sync() override { - int tmp = 0; - if (is_server) { - send_data_internal(&tmp, 1); - recv_data_internal(&tmp, 1); - } else { - recv_data_internal(&tmp, 1); - send_data_internal(&tmp, 1); - flush_unlocked(); - } - } - void send_data_internal(const void *data, int64_t len) override { expecting(len >= 0, "NetIO::send_data: negative len"); if (len == 0) return; @@ -237,7 +224,9 @@ class NetIO : public IOChannel { public: if (!send_dirty) return; ++flushes_count; if (send_ptr) { send_raw(send_buf, send_ptr); send_ptr = 0; } - fflush(stream); + expecting(::fflush(stream) == 0, [&] { + return std::string("NetIO: fflush failed: ") + std::strerror(errno); + }); send_dirty = false; } diff --git a/emp-tool/runtime/io/tls_io_channel.h b/emp-tool/runtime/io/tls_io_channel.h index ea8d9bc..c17965e 100644 --- a/emp-tool/runtime/io/tls_io_channel.h +++ b/emp-tool/runtime/io/tls_io_channel.h @@ -101,7 +101,7 @@ inline void install_sigpipe_ignore_once() { class TLSIO : public IOChannel { public: int sock = -1; - bool is_server, quiet; + bool quiet; bool is_tls_server; SSL_CTX *ctx = nullptr; @@ -130,11 +130,11 @@ class TLSIO : public IOChannel { public: expecting(port >= 0 && port <= 65535, "TLSIO: invalid port number"); tls_detail::install_sigpipe_ignore_once(); - is_server = (address == nullptr); + const bool is_listener = (address == nullptr); is_tls_server = cfg.is_tls_server; - init_from_sock(is_server ? tcp::server_listen(port, cfg.socket_options) - : tcp::client_connect(address, port, - cfg.socket_options), + init_from_sock(is_listener ? tcp::server_listen(port, cfg.socket_options) + : tcp::client_connect(address, port, + cfg.socket_options), cfg); if (!quiet) std::cout << "TLS connected\n"; } @@ -159,7 +159,6 @@ class TLSIO : public IOChannel { public: cfg.socket_options.receive_buffer_size == 0, "TLSIO: socket buffer options cannot be applied to an already-connected socket"); tls_detail::install_sigpipe_ignore_once(); - is_server = false; init_from_sock(existing_sock, cfg); } @@ -284,19 +283,6 @@ class TLSIO : public IOChannel { public: } } - // 1-byte ping/pong handshake to verify both directions are alive. - // is_server (the TCP-acceptor bit) decides who sends first. - void sync() override { - int tmp = 0; - if (is_server) { - send_data_internal(&tmp, 1); - recv_data_internal(&tmp, 1); - } else { - recv_data_internal(&tmp, 1); - send_data_internal(&tmp, 1); - } - } - void send_data_internal(const void *data, int64_t len) override { expecting(len >= 0, "TLSIO::send_data: negative len"); if (len == 0) return; diff --git a/emp-tool/runtime/io/trace_io.h b/emp-tool/runtime/io/trace_io.h index f871c47..66e1db8 100644 --- a/emp-tool/runtime/io/trace_io.h +++ b/emp-tool/runtime/io/trace_io.h @@ -23,10 +23,14 @@ // deterministic seeds. Test-mode-off → traces are non-reproducible. #include "emp-tool/runtime/io/io_channel.h" +#include #include #include #include +#include #include +#include +#include namespace emp { @@ -34,19 +38,14 @@ class TraceIO : public IOChannel { public: // `under` is borrowed (not owned). `prefix` selects the trace file // names: ".send" for outbound bytes, ".recv" for - // inbound. Files are opened binary-write, truncating. + // inbound. Files are opened binary-write, truncating, with mode 0600. TraceIO(IOChannel* under, const std::string& prefix) : under_(under) { + expecting(under_ != nullptr, "TraceIO: underlying channel is null"); const std::string send_path = prefix + ".send"; const std::string recv_path = prefix + ".recv"; - send_fp_ = std::fopen(send_path.c_str(), "wb"); - expecting(send_fp_ != nullptr, [&] { - return "TraceIO: cannot open " + send_path + " for write"; - }); - recv_fp_ = std::fopen(recv_path.c_str(), "wb"); - expecting(recv_fp_ != nullptr, [&] { - return "TraceIO: cannot open " + recv_path + " for write"; - }); + send_fp_ = open_trace_file(send_path); + recv_fp_ = open_trace_file(recv_path); } ~TraceIO() override { @@ -58,8 +57,7 @@ class TraceIO : public IOChannel { expecting(nbyte >= 0, "TraceIO::send_data_internal: negative byte count"); if (nbyte == 0) return; - // Tee first, deliver after, so a crash mid-write still leaves - // a trace prefix that matches what the peer didn't yet see. + // Record outbound bytes before delivering them to the transport. const size_t bytes = static_cast(nbyte); expecting(std::fwrite(data, 1, bytes, send_fp_) == bytes, "TraceIO: short write to .send"); @@ -78,10 +76,39 @@ class TraceIO : public IOChannel { "TraceIO: short write to .recv"); } - void flush() override { under_->flush(); } - void sync() override { under_->sync(); } + void flush() override { + flush_trace_file(send_fp_, ".send"); + flush_trace_file(recv_fp_, ".recv"); + under_->flush(); + } private: + static std::FILE* open_trace_file(const std::string& path) { + int fd; + do { + fd = ::open(path.c_str(), O_WRONLY | O_CREAT | O_TRUNC, 0600); + } while (fd < 0 && errno == EINTR); + expecting(fd >= 0, [&] { + return "TraceIO: cannot open " + path + ": " + std::strerror(errno); + }); + expecting(::fchmod(fd, S_IRUSR | S_IWUSR) == 0, [&] { + return "TraceIO: cannot secure " + path + ": " + std::strerror(errno); + }); + std::FILE* file = ::fdopen(fd, "wb"); + expecting(file != nullptr, [&] { + return "TraceIO: fdopen failed for " + path + ": " + + std::strerror(errno); + }); + return file; + } + + static void flush_trace_file(std::FILE* file, const char* suffix) { + expecting(std::fflush(file) == 0, [&] { + return std::string("TraceIO: flush ") + suffix + " failed: " + + std::strerror(errno); + }); + } + IOChannel* under_; std::FILE* send_fp_ = nullptr; std::FILE* recv_fp_ = nullptr; diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index e5d08c7..a94e852 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -82,6 +82,7 @@ add_test_case(runtime test_ecc) add_test_case(runtime test_ro) add_test_case_with_run(runtime test_netio) add_test_case_with_run(runtime test_tlsio) +add_test_case(runtime test_traceio) set_tests_properties(test_netio test_tlsio PROPERTIES TIMEOUT 130) add_test_case(runtime test_halfgate) # half-gate garble/eval free functions diff --git a/test/runtime/test_netio.cpp b/test/runtime/test_netio.cpp index f7c56cb..4807d37 100644 --- a/test/runtime/test_netio.cpp +++ b/test/runtime/test_netio.cpp @@ -7,7 +7,6 @@ // send_block / recv_block block-typed wrapper // send_bool / recv_bool packed via bools_to_bits // flush() drain outbound only (no peer coupling) -// sync() 1-byte ping/pong handshake // make_sibling() more connections on the same port // tcp::SocketOptions pre-handshake socket buffer sizing // SocketOptions::for_bandwidth_and_rtt @@ -16,6 +15,7 @@ // Test functions below are templated on the IO type so correctness and // regression checks can be reused by IO implementations. +#include #include #include #include @@ -44,6 +44,22 @@ static bool dies(F &&f) { return !(WIFEXITED(status) && WEXITSTATUS(status) == 0); } +static void round_trip_marker(NetIO &io, int party) { + const char marker = 0x5a; + char received = 0; + if (party == ALICE) { + io.send_data(&marker, 1); + io.flush(); + io.recv_data(&received, 1); + } else { + io.recv_data(&received, 1); + io.send_data(&marker, 1); + io.flush(); + } + expecting(received == marker, + "NetIO test: sibling round-trip marker mismatch"); +} + // ------------------------------------------------------------------------- // run_correctness(): byte stream round-trip at unaligned offsets, then bool // packing round-trip at unaligned bool offsets. Each side asserts on the @@ -217,7 +233,7 @@ static void run_sibling_regression(int port, int party) { for (int round = 0; round < 16; ++round) { NetIO *source = round < 8 ? primary.get() : siblings.front().get(); siblings.push_back(source->make_sibling()); - siblings.back()->sync(); + round_trip_marker(*siblings.back(), party); if (round == 7) primary.reset(); } } @@ -225,7 +241,7 @@ static void run_sibling_regression(int port, int party) { auto primary = party == ALICE ? NetIO::listen(port, true) : NetIO::connect(peer_ip(), port, true); auto sibling = primary->make_sibling(); - sibling->sync(); + round_trip_marker(*sibling, party); } if (party == ALICE) cout << "NetIO shared-listener regression: OK\n"; } @@ -286,6 +302,20 @@ static void run_socket_options_factory_regression() { }), "NetIO test: unavailable socket buffer was not rejected"); } +static void run_flush_failure_regression() { + expecting(dies([] { + int sockets[2]; + expecting(::socketpair(AF_UNIX, SOCK_STREAM, 0, sockets) == 0, + "NetIO test: socketpair failed"); + ::signal(SIGPIPE, SIG_IGN); + ::close(sockets[1]); + NetIO io(sockets[0], true); + const char byte = 1; + io.send_data(&byte, 1); + io.flush(); + }), "NetIO test: fflush failure was ignored"); +} + static void run_socket_options_regression(int port, int party) { tcp::SocketOptions options; options.send_buffer_size = 128 * 1024; @@ -304,12 +334,12 @@ static void run_socket_options_regression(int port, int party) { } auto sibling = primary->make_sibling(); - sibling->sync(); + round_trip_marker(*sibling, party); expect_socket_options(*sibling, options); primary.reset(); auto next = sibling->make_sibling(); - next->sync(); + round_trip_marker(*next, party); expect_socket_options(*next, options); if (party == ALICE) cout << "NetIO socket-options regression: OK\n"; } @@ -347,6 +377,7 @@ int main(int argc, char **argv) { party = parse_party(argv); port = peer_port(); run_socket_options_factory_regression(); + run_flush_failure_regression(); // Five contiguous ports: main, two send-only cases, sibling regression, // socket-options regression. diff --git a/test/runtime/test_tlsio.cpp b/test/runtime/test_tlsio.cpp index cb19132..7463b18 100644 --- a/test/runtime/test_tlsio.cpp +++ b/test/runtime/test_tlsio.cpp @@ -7,7 +7,6 @@ // send_block / recv_block block-typed wrapper (inherited) // send_bool / recv_bool packed via bools_to_bits (inherited) // flush() drain outbound coalescing buffer -// sync() 1-byte ping/pong handshake // TLSConfig::socket_options pre-handshake socket buffer sizing // // Same flush contract and thread-safety rules as NetIO. The test diff --git a/test/runtime/test_traceio.cpp b/test/runtime/test_traceio.cpp new file mode 100644 index 0000000..33c6377 --- /dev/null +++ b/test/runtime/test_traceio.cpp @@ -0,0 +1,127 @@ +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include + +#include "emp-tool/emp-tool.h" + +using namespace emp; + +template +static bool dies(F&& f) { + pid_t pid = ::fork(); + expecting(pid >= 0, "TraceIO test: fork failed"); + if (pid == 0) { + std::freopen("/dev/null", "w", stderr); + f(); + ::_exit(0); + } + int status = 0; + ::waitpid(pid, &status, 0); + return !(WIFEXITED(status) && WEXITSTATUS(status) == 0); +} + +class MemoryIO final : public IOChannel { +public: + explicit MemoryIO(std::vector incoming) + : incoming_(std::move(incoming)) {} + + void send_data_internal(const void* data, int64_t nbyte) override { + const auto* bytes = static_cast(data); + sent.insert(sent.end(), bytes, bytes + nbyte); + } + + void recv_data_internal(void* data, int64_t nbyte) override { + expecting(nbyte >= 0 && offset_ + static_cast(nbyte) <= incoming_.size(), + "TraceIO test: MemoryIO input exhausted"); + std::memcpy(data, incoming_.data() + offset_, static_cast(nbyte)); + offset_ += static_cast(nbyte); + } + + void flush() override { ++flush_count; } + + std::vector sent; + int flush_count = 0; + +private: + std::vector incoming_; + size_t offset_ = 0; +}; + +static std::vector read_file(const std::string& path) { + std::ifstream input(path, std::ios::binary); + expecting(input.good(), "TraceIO test: cannot read trace file"); + return std::vector(std::istreambuf_iterator(input), + std::istreambuf_iterator()); +} + +static void make_permissive_file(const std::string& path) { + int fd = ::open(path.c_str(), O_WRONLY | O_CREAT | O_TRUNC, 0666); + expecting(fd >= 0, "TraceIO test: cannot create existing trace file"); + expecting(::fchmod(fd, 0666) == 0, + "TraceIO test: cannot set existing trace permissions"); + expecting(::close(fd) == 0, + "TraceIO test: cannot close existing trace file"); +} + +static void expect_private_file(const std::string& path) { + struct stat info; + expecting(::stat(path.c_str(), &info) == 0, + "TraceIO test: cannot stat trace file"); + expecting((info.st_mode & 0777) == 0600, + "TraceIO test: trace file permissions are not 0600"); +} + +int main() { + char directory_template[] = "/tmp/emp_traceio_test.XXXXXX"; + char* directory = ::mkdtemp(directory_template); + expecting(directory != nullptr, + "TraceIO test: cannot create temporary directory"); + const std::string prefix = std::string(directory) + "/trace"; + const std::string send_path = prefix + ".send"; + const std::string recv_path = prefix + ".recv"; + + expecting(dies([&] { TraceIO trace(nullptr, prefix + ".null"); }), + "TraceIO test: null underlying channel was accepted"); + + make_permissive_file(send_path); + make_permissive_file(recv_path); + MemoryIO under({'r', 'e', 'c', 'v'}); + { + TraceIO trace(&under, prefix); + const char outbound[] = {'s', 'e', 'n', 'd'}; + char inbound[sizeof(outbound)] = {}; + trace.send_data(outbound, sizeof(outbound)); + trace.recv_data(inbound, sizeof(inbound)); + trace.flush(); + + expecting(under.sent == std::vector(outbound, outbound + sizeof(outbound)), + "TraceIO test: outbound delegation mismatch"); + expecting(std::memcmp(inbound, "recv", sizeof(inbound)) == 0, + "TraceIO test: inbound delegation mismatch"); + expecting(read_file(send_path) == under.sent, + "TraceIO test: .send was not flushed"); + expecting(read_file(recv_path) == std::vector(inbound, inbound + sizeof(inbound)), + "TraceIO test: .recv was not flushed"); + expecting(under.flush_count == 1, + "TraceIO test: underlying flush was not delegated"); + expect_private_file(send_path); + expect_private_file(recv_path); + } + + expecting(::unlink(send_path.c_str()) == 0, + "TraceIO test: cannot remove .send"); + expecting(::unlink(recv_path.c_str()) == 0, + "TraceIO test: cannot remove .recv"); + expecting(::rmdir(directory) == 0, + "TraceIO test: cannot remove temporary directory"); +} From 8758e2c2d9084b6aa05f34dc2a0ad901ceb9a3df Mon Sep 17 00:00:00 2001 From: Xiao Wang Date: Sat, 22 Aug 2026 08:25:01 -0500 Subject: [PATCH 4/7] Enforce runtime fail-stop and test-mode contracts Flush fatal diagnostics before terminating. Require exact deterministic-test activation, reserve lane zero for the owner thread, derive child lanes deterministically, and reject invalid lane and ThreadPool configurations. Add focused fail-stop, deterministic threading, and pool-lifecycle coverage. --- docs/test_mode.md | 18 ++-- emp-tool/runtime/core/error.h | 1 + emp-tool/runtime/core/test_mode.h | 59 +++++++------ emp-tool/third_party/ThreadPool.h | 1 + test/CMakeLists.txt | 4 +- test/runtime/test_error.cpp | 14 +++- test/runtime/test_test_mode.cpp | 134 +++++++++++++++++++++++++++++- test/runtime/test_thread_pool.cpp | 74 +++++++++++++++++ 8 files changed, 268 insertions(+), 37 deletions(-) create mode 100644 test/runtime/test_thread_pool.cpp diff --git a/docs/test_mode.md b/docs/test_mode.md index df6498e..71b4aa7 100644 --- a/docs/test_mode.md +++ b/docs/test_mode.md @@ -56,7 +56,8 @@ emp::set_test_mode(true); // before any PRG() default-construction ``` The env var is read once at first call to `is_test_mode()` and -cached. `set_test_mode()` overrides it programmatically. +cached. Only the exact value `1` enables test mode; other values leave it off. +`set_test_mode()` overrides it programmatically. The first activation by either mechanism prints a prominent warning to `stderr`, once per process, that default PRG seeds and EC scalar randomness are @@ -64,9 +65,9 @@ deterministic and insecure. The warning happens at activation rather than on each random draw, so it adds no work to the randomness hot path. Never process real secrets in a process running in test mode. -`reset_test_seed_counter()` rewinds every lane's ordinal and -releases lane 0 — call it between independent test iterations to -get reproducible PRG sequences within one process. +`reset_test_seed_counter()` rewinds every lane's ordinal and releases lane 0. +Call it between independent test iterations, after joining threads and draining +pool futures, to get reproducible PRG sequences within one process. ## Multi-threading: lanes @@ -88,9 +89,12 @@ order is deterministic), never discovered by the worker itself. }); ``` -- **Forgetting is loud.** A second thread drawing from lane 0 would - replay the main thread's streams byte-for-byte — silently wrong — - so test mode aborts with a pointer to this document instead. + Manually assigned lane ids must be nonzero and unique among concurrently + active work. Lane 0 is reserved for the unscoped main thread. + +- **Forgetting is loud.** A second thread drawing from lane 0 or deriving a + child lane from it would replay deterministic streams byte-for-byte, so test + mode aborts with a pointer to this document instead. Lanes make the *randomness* reproducible. Byte-identical *traces* additionally require that each traced channel has a single writer diff --git a/emp-tool/runtime/core/error.h b/emp-tool/runtime/core/error.h index 16dca63..b819335 100644 --- a/emp-tool/runtime/core/error.h +++ b/emp-tool/runtime/core/error.h @@ -24,6 +24,7 @@ inline void error(const char *s, int line = __builtin_LINE(), const char *file = __builtin_FILE()) { std::fprintf(stderr, "%s at %s:%d\n", s, file, line); + std::fflush(stderr); // _Exit, not exit(): error() can fire from a worker thread while sibling // workers still own heap state. Running destructors/atexit handlers in that // situation races their live work; terminate the process immediately. diff --git a/emp-tool/runtime/core/test_mode.h b/emp-tool/runtime/core/test_mode.h index 9d9f505..299fb1c 100644 --- a/emp-tool/runtime/core/test_mode.h +++ b/emp-tool/runtime/core/test_mode.h @@ -18,9 +18,10 @@ // depend only on per-lane program order, never on cross-thread // scheduling, so multi-threaded runs reproduce. ThreadPool::enqueue // derives and installs a lane per task automatically; hand-spawned -// threads wrap their body in test_lane_scope. A second thread drawing -// from lane 0 aborts: two threads sharing a lane would replay identical -// "random" streams — silently wrong rather than merely nondeterministic. +// threads wrap their body in test_lane_scope. A second thread using lane +// 0 to draw a seed or derive a child lane aborts: two threads sharing a +// lane would replay identical "random" streams — silently wrong rather +// than merely nondeterministic. #include "emp-tool/runtime/core/error.h" @@ -28,6 +29,7 @@ #include #include #include +#include #include #include @@ -58,7 +60,7 @@ inline std::atomic& test_mode_flag() { static std::atomic flag( []() { const char* v = std::getenv("EMP_TEST_MODE"); - const bool enabled = v != nullptr && v[0] == '1'; + const bool enabled = v != nullptr && std::strcmp(v, "1") == 0; if (enabled) warn_insecure_test_mode_once(); return enabled; }()); @@ -99,16 +101,31 @@ inline void sync_test_epoch(TestSeedTls& s) { } } -// Owner token of lane 0: the one thread allowed to draw main-lane -// seeds. Cleared by reset_test_seed_counter(), so sequential +// Owner token of lane 0: the one thread allowed to draw main-lane seeds or +// derive child lanes. Cleared by reset_test_seed_counter(), so sequential // independent units may run on different threads. inline std::atomic& lane0_owner() { static std::atomic owner(0); return owner; } +inline std::atomic& next_thread_token() { + static std::atomic next(1); + return next; +} inline uint64_t this_thread_token() { - // Nonzero hash of the thread id; 0 is the "unowned" sentinel. - return (uint64_t)std::hash()(std::this_thread::get_id()) | 1ULL; + thread_local const uint64_t token = + next_thread_token().fetch_add(1, std::memory_order_relaxed); + expecting(token != 0, "test mode: thread token space exhausted"); + return token; +} +inline void claim_lane0() { + auto& owner = lane0_owner(); + const uint64_t token = this_thread_token(); + uint64_t expected = 0; + expecting(owner.compare_exchange_strong(expected, token) || + expected == token, + "test mode: a second thread used lane 0; run spawned work " + "under emp::test_lane_scope (see docs/test_mode.md)"); } // splitmix64 finalizer: full-avalanche 64-bit mix for deriving child @@ -149,19 +166,7 @@ struct TestSeed { inline TestSeed next_test_seed() { auto& s = detail::test_seed_tls(); detail::sync_test_epoch(s); - if (s.lane == 0) { - // Only one thread may consume main-lane seeds; a second one - // would replay the same streams. Always-on: test mode usually - // runs under Release/NDEBUG builds. - auto& owner = detail::lane0_owner(); - const uint64_t token = detail::this_thread_token(); - uint64_t expected = 0; - expecting(owner.compare_exchange_strong(expected, token) || - expected == token, - "test mode: a second thread drew lane-0 randomness; run " - "spawned work under emp::test_lane_scope (see " - "docs/test_mode.md)"); - } + if (s.lane == 0) detail::claim_lane0(); return {s.lane, s.ctr++}; } @@ -174,6 +179,7 @@ inline TestSeed next_test_seed() { inline uint64_t next_test_child_lane() { auto& s = detail::test_seed_tls(); detail::sync_test_epoch(s); + if (s.lane == 0) detail::claim_lane0(); uint64_t lane = detail::mix64(detail::mix64(s.lane) ^ s.child_ctr++); if (lane == 0) lane = 1; // 0 is reserved for the main thread return lane; @@ -186,6 +192,8 @@ inline uint64_t next_test_child_lane() { class test_lane_scope { public: explicit test_lane_scope(uint64_t lane) : saved_(detail::test_seed_tls()) { + expecting(lane != 0, + "test_lane_scope: lane 0 is reserved for the main thread"); auto& s = detail::test_seed_tls(); s.lane = lane; s.ctr = 0; @@ -207,10 +215,11 @@ inline uint64_t current_test_seed_epoch() { return detail::test_seed_epoch().load(); } -// Rewind every lane's draw ordinal (lazily, when each thread next -// draws) and release lane 0. Use before each independent unit (e.g. -// each protocol in a trace) to make that unit's randomness -- and thus -// its wire bytes -- independent of whatever consumed seeds before it. +// Rewind every lane's draw ordinal (lazily, when each thread next draws) and +// release lane 0. All work that can draw seeds or derive child lanes must be +// quiescent first. Use before each independent unit (e.g. each protocol in a +// trace) to make that unit's randomness -- and thus its wire bytes -- +// independent of whatever consumed seeds before it. inline void reset_test_seed_counter() { detail::test_seed_epoch().fetch_add(1); detail::lane0_owner().store(0); diff --git a/emp-tool/third_party/ThreadPool.h b/emp-tool/third_party/ThreadPool.h index 37df84f..881a030 100644 --- a/emp-tool/third_party/ThreadPool.h +++ b/emp-tool/third_party/ThreadPool.h @@ -72,6 +72,7 @@ inline size_t ThreadPool::size() const { return workers.size(); } // the constructor just launches some amount of workers inline ThreadPool::ThreadPool(size_t threads) : stop(false) { + emp::expecting(threads > 0, "ThreadPool: worker count must be positive"); for (size_t i = 0; i < threads; ++i) workers.emplace_back([this] { for (;;) { diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index a94e852..653b51e 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -72,7 +72,9 @@ add_test_case(runtime test_prp) add_test_case(runtime test_test_mode) set_tests_properties(test_test_mode PROPERTIES ENVIRONMENT "EMP_TEST_MODE=1" - PASS_REGULAR_EXPRESSION "WARNING: EMP TEST MODE ENABLED - RANDOMNESS IS DETERMINISTIC AND INSECURE") + PASS_REGULAR_EXPRESSION "WARNING: EMP TEST MODE ENABLED - RANDOMNESS IS DETERMINISTIC AND INSECURE" + FAIL_REGULAR_EXPRESSION "FAIL;CORRECTNESS FAILURE") +add_test_case(runtime test_thread_pool) add_test_case(runtime test_utils) # Exception-free public surface (docs/api_conventions.md): compiled with # -fno-exceptions, so any `throw` reachable from the umbrella fails the build. diff --git a/test/runtime/test_error.cpp b/test/runtime/test_error.cpp index ae41625..cb89c5b 100644 --- a/test/runtime/test_error.cpp +++ b/test/runtime/test_error.cpp @@ -69,9 +69,20 @@ static bool check_lazy_message_stays_lazy() { return !built; } +static bool check_buffered_diagnostic_is_flushed() { + DeathResult result = run_child([] { + char buffer[4096]; + if (setvbuf(stderr, buffer, _IOFBF, sizeof(buffer)) != 0) _exit(2); + expecting(false, "buffered fatal diagnostic"); + }); + return result.died && + result.stderr_text.find("buffered fatal diagnostic") != string::npos; +} + static bool run_correctness() { bool once = check_expecting_evaluates_once(); bool lazy = check_lazy_message_stays_lazy(); + bool buffered = check_buffered_diagnostic_is_flushed(); DeathResult expectation = run_child([] { expecting(false, "expected expectation failure"); }); DeathResult dynamic = run_child([] { @@ -83,11 +94,12 @@ static bool run_correctness() { bool dynamic_message = dynamic.stderr_text.find("expected lazy failure") != string::npos; cout << " expecting evaluates once " << (once ? "OK" : "FAIL") << "\n"; cout << " lazy message stays lazy " << (lazy ? "OK" : "FAIL") << "\n"; + cout << " buffered diagnostic is flushed " << (buffered ? "OK" : "FAIL") << "\n"; cout << " failed expectation terminates " << (expectation.died ? "OK" : "FAIL") << "\n"; cout << " expectation reports its message " << (expectation_message ? "OK" : "FAIL") << "\n"; cout << " expectation reports caller site " << (caller_location ? "OK" : "FAIL") << "\n"; cout << " lazy failure reports/terminates " << (dynamic.died && dynamic_message ? "OK" : "FAIL") << "\n"; - return once && lazy && expectation.died && expectation_message && caller_location && + return once && lazy && buffered && expectation.died && expectation_message && caller_location && dynamic.died && dynamic_message; } diff --git a/test/runtime/test_test_mode.cpp b/test/runtime/test_test_mode.cpp index acf7a97..9c623f9 100644 --- a/test/runtime/test_test_mode.cpp +++ b/test/runtime/test_test_mode.cpp @@ -16,12 +16,60 @@ #include "emp-tool/emp-tool.h" +#include +#include #include +#include +#include +#include #include using namespace emp; using namespace std; +template +static int run_child(F&& f) { + pid_t pid = fork(); + if (pid == 0) { + close(STDOUT_FILENO); + close(STDERR_FILENO); + f(); + _exit(0); + } + if (pid < 0) return -1; + int status = 0; + if (waitpid(pid, &status, 0) != pid) return -1; + return status; +} + +static bool child_died(int status) { + return status >= 0 && + (!WIFEXITED(status) || WEXITSTATUS(status) != 0); +} + +static bool child_succeeded(int status) { + return status >= 0 && WIFEXITED(status) && WEXITSTATUS(status) == 0; +} + +static bool probe_environment(const char* executable, const char* value, + bool expected) { + pid_t pid = fork(); + if (pid == 0) { + int env_result = value == nullptr + ? unsetenv("EMP_TEST_MODE") + : setenv("EMP_TEST_MODE", value, 1); + if (env_result != 0) _exit(127); + close(STDOUT_FILENO); + close(STDERR_FILENO); + execlp(executable, executable, "--probe-env", + expected ? "on" : "off", static_cast(nullptr)); + _exit(127); + } + if (pid < 0) return false; + int status = 0; + return waitpid(pid, &status, 0) == pid && child_succeeded(status); +} + // ---------- example ---------- @@ -183,22 +231,102 @@ static bool check_child_lane_derivation_deterministic() { return ok; } -static bool run_correctness() { +static bool check_environment_contract(const char* executable) { + bool ok = probe_environment(executable, "1", true) && + probe_environment(executable, nullptr, false) && + probe_environment(executable, "", false) && + probe_environment(executable, "0", false) && + probe_environment(executable, "10", false) && + probe_environment(executable, "1x", false) && + probe_environment(executable, "true", false); + cout << " [EMP_TEST_MODE accepts exactly 1] " << (ok ? "OK" : "FAIL") << "\n"; + return ok; +} + +static bool check_lane0_seed_owner() { + bool ok = child_died(run_child([] { + reset_test_seed_counter(); + (void)next_test_seed(); + thread second([] { (void)next_test_seed(); }); + second.join(); + })); + cout << " [second lane-0 seed thread is rejected] " << (ok ? "OK" : "FAIL") << "\n"; + return ok; +} + +static bool check_lane0_child_owner() { + bool ok = child_died(run_child([] { + reset_test_seed_counter(); + (void)next_test_child_lane(); + thread second([] { (void)next_test_child_lane(); }); + second.join(); + })); + cout << " [second lane-0 child creator is rejected] " << (ok ? "OK" : "FAIL") << "\n"; + return ok; +} + +static bool check_lane0_thread_lifetime() { + bool ok = child_died(run_child([] { + reset_test_seed_counter(); + thread first([] { (void)next_test_seed(); }); + first.join(); + thread second([] { (void)next_test_seed(); }); + second.join(); + })); + cout << " [lane-0 ownership spans thread lifetimes] " << (ok ? "OK" : "FAIL") << "\n"; + return ok; +} + +static bool check_lane_zero_rejected() { + bool ok = child_died(run_child([] { test_lane_scope scope(0); })); + cout << " [test_lane_scope rejects lane 0] " << (ok ? "OK" : "FAIL") << "\n"; + return ok; +} + +static bool check_reset_allows_lane0_handoff() { + bool ok = child_succeeded(run_child([] { + reset_test_seed_counter(); + TestSeed first_seed{}; + thread first([&] { first_seed = next_test_seed(); }); + first.join(); + if (first_seed.lane != 0 || first_seed.ordinal != 0) _exit(2); + + reset_test_seed_counter(); + TestSeed second_seed{}; + thread second([&] { second_seed = next_test_seed(); }); + second.join(); + if (second_seed.lane != 0 || second_seed.ordinal != 0) _exit(3); + })); + cout << " [reset permits quiescent lane-0 handoff] " << (ok ? "OK" : "FAIL") << "\n"; + return ok; +} + +static bool run_correctness(const char* executable) { cout << "=== correctness ===\n"; bool ok = true; + ok &= check_environment_contract(executable); ok &= check_main_lane_sequence(); ok &= check_lane_scope_nesting(); + ok &= check_lane0_seed_owner(); + ok &= check_lane0_child_owner(); + ok &= check_lane0_thread_lifetime(); + ok &= check_lane_zero_rejected(); + ok &= check_reset_allows_lane0_handoff(); ok &= check_pool_determinism(); ok &= check_enqueue_order_is_the_lane(); ok &= check_child_lane_derivation_deterministic(); return ok; } -int main(int /*argc*/, char ** /*argv*/) { +int main(int argc, char** argv) { + if (argc == 3 && std::strcmp(argv[1], "--probe-env") == 0) { + bool expected = std::strcmp(argv[2], "on") == 0; + return is_test_mode() == expected ? 0 : 1; + } set_test_mode(true); example(); cout << "\n"; - if (!run_correctness()) { + if (!run_correctness(argv[0])) { cerr << "CORRECTNESS FAILURE\n"; return 1; } diff --git a/test/runtime/test_thread_pool.cpp b/test/runtime/test_thread_pool.cpp new file mode 100644 index 0000000..87a6137 --- /dev/null +++ b/test/runtime/test_thread_pool.cpp @@ -0,0 +1,74 @@ +// third_party/ThreadPool.h — fixed-size worker pool with future-returning tasks. +// Read example() first; the rest is verification. +// +// What's in ThreadPool.h: +// ThreadPool(n) start n workers +// enqueue(f, args...) schedule one task and return its future +// size() report the fixed worker count + +#include "emp-tool/third_party/ThreadPool.h" + +#include +#include +#include +#include +#include +#include + +using namespace std; + +template +static bool dies(F&& f) { + pid_t pid = fork(); + if (pid == 0) { + close(STDERR_FILENO); + f(); + _exit(0); + } + if (pid < 0) return false; + int status = 0; + if (waitpid(pid, &status, 0) != pid) return false; + return !WIFEXITED(status) || WEXITSTATUS(status) != 0; +} + +static void example() { + ThreadPool pool(2); + auto answer = pool.enqueue([] { return 6 * 7; }); + cout << "ThreadPool task result = " << answer.get() << "\n"; +} + +static bool check_tasks_complete() { + ThreadPool pool(3); + vector> results; + for (int i = 0; i < 32; ++i) + results.push_back(pool.enqueue([i] { return i * i; })); + for (int i = 0; i < 32; ++i) + if (results[i].get() != i * i) return false; + return pool.size() == 3; +} + +static bool check_destructor_drains_queue() { + atomic completed{0}; + { + ThreadPool pool(2); + for (int i = 0; i < 32; ++i) + pool.enqueue([&completed] { completed.fetch_add(1); }); + } + return completed.load() == 32; +} + +static bool run_correctness() { + bool tasks = check_tasks_complete(); + bool drains = check_destructor_drains_queue(); + bool rejects_zero = dies([] { ThreadPool pool(0); }); + cout << " tasks complete and size is fixed " << (tasks ? "OK" : "FAIL") << "\n"; + cout << " destructor drains queued work " << (drains ? "OK" : "FAIL") << "\n"; + cout << " zero workers are rejected " << (rejects_zero ? "OK" : "FAIL") << "\n"; + return tasks && drains && rejects_zero; +} + +int main() { + example(); + cout << "\n=== correctness ===\n"; + return run_correctness() ? 0 : 1; +} From 30e9acf1c1ebf5bf52de01aa953baaa736f02333 Mon Sep 17 00:00:00 2001 From: Xiao Wang Date: Sat, 22 Aug 2026 08:25:01 -0500 Subject: [PATCH 5/7] Define zero-length block helper behavior Accept null pointers only when a block-array operation has length zero and reject negative lengths before pointer arithmetic. Add direct zero-length and invalid-length coverage. --- emp-tool/runtime/core/block.hpp | 2 ++ test/runtime/test_block.cpp | 14 ++++++++++++++ 2 files changed, 16 insertions(+) diff --git a/emp-tool/runtime/core/block.hpp b/emp-tool/runtime/core/block.hpp index 39a56a0..1a32172 100644 --- a/emp-tool/runtime/core/block.hpp +++ b/emp-tool/runtime/core/block.hpp @@ -66,6 +66,8 @@ inline block set_bit(const block & a, int i) { } inline std::string to_hex(const void* data, size_t n) { + expecting(n <= std::string{}.max_size() / 2, + "to_hex: input too large"); static const char digits[] = "0123456789abcdef"; const unsigned char* b = static_cast(data); std::string s(2 * n, '0'); diff --git a/test/runtime/test_block.cpp b/test/runtime/test_block.cpp index 9c21feb..947af1d 100644 --- a/test/runtime/test_block.cpp +++ b/test/runtime/test_block.cpp @@ -6,6 +6,7 @@ // block, makeBlock(hi, lo), zero_block, all_one_block // getLSB(b), set_bit(b, i) // sigma(b) linear orthomorphism (Guo et al.) +// to_hex(data, n) lowercase byte-order hex // xorBlocks_arr(res, x, y, n) element-wise XOR // xorBlocks_arr(res, x, y_block, n) broadcast XOR // xorBlocksTo_arr(dst, src, n) in-place dst[i] ^= src[i] @@ -175,6 +176,15 @@ static bool check_xorBlocks_arr() { return true; } +static bool check_zero_length_helpers() { + block* out = nullptr; + const block* in = nullptr; + xorBlocks_arr(out, in, in, 0); + xorBlocksTo_arr(out, in, 0); + xorBlocks_arr(out, in, zero_block, 0); + return cmpBlock(in, in, 0) && to_hex(nullptr, 0).empty(); +} + static bool check_cmpBlock() { PRG prg; for (int sz : {1, 4, 33, 257}) { @@ -201,6 +211,9 @@ static bool check_invalid_ranges_rejected() { && dies([&] { xorBlocksTo_arr(&out, &a, -1); }) && dies([&] { xorBlocks_arr(&out, &a, b, -1); }) && dies([&] { (void)cmpBlock(&a, &b, -1); }) + && dies([&] { + (void)to_hex(nullptr, std::string{}.max_size() / 2 + 1); + }) && dies([&] { bools_to_bits(packed, bits, -1); }) && dies([&] { bits_to_bools(bits, packed, -1); }) && dies([&] { sse_trans(packed, packed, 7, 8); }) @@ -333,6 +346,7 @@ static bool run_correctness() { {"set_bit", check_set_bit}, {"sigma linear + formula", check_sigma_linear_and_formula}, {"xorBlocks_arr / xorBlocksTo", check_xorBlocks_arr}, + {"zero-length block helpers", check_zero_length_helpers}, {"cmpBlock", check_cmpBlock}, {"invalid ranges rejected", check_invalid_ranges_rejected}, {"sse_trans round-trip", check_sse_trans_roundtrip}, From bc6f7d90629d34812993410e1083523b3cae381e Mon Sep 17 00:00:00 2001 From: Xiao Wang Date: Sat, 22 Aug 2026 08:25:01 -0500 Subject: [PATCH 6/7] Validate core utilities and transpose remainders Strictly parse party and port inputs, preserve the steady-clock type through timing helpers, and validate bit-packing lengths. Fix the portable SIMD include bootstrap and handle every remainder row in the scalar transpose tail. Add independent scalar-reference and utility boundary tests. --- emp-tool/runtime/core/simd_tier.h | 3 +- emp-tool/runtime/core/transpose.hpp | 59 +++++++------------ emp-tool/runtime/core/utils.h | 6 +- emp-tool/runtime/core/utils.hpp | 45 ++++++++++---- emp-tool/third_party/ThreadPool.h | 13 ++++- test/runtime/test_block.cpp | 48 +++++++++++++++ test/runtime/test_thread_pool.cpp | 31 +++++++++- test/runtime/test_utils.cpp | 91 ++++++++++++++++++++++++++++- 8 files changed, 239 insertions(+), 57 deletions(-) diff --git a/emp-tool/runtime/core/simd_tier.h b/emp-tool/runtime/core/simd_tier.h index 9172cb1..181e80d 100644 --- a/emp-tool/runtime/core/simd_tier.h +++ b/emp-tool/runtime/core/simd_tier.h @@ -1,7 +1,6 @@ #ifndef EMP_SIMD_TIER_H__ #define EMP_SIMD_TIER_H__ -#include "emp-tool/runtime/core/block.h" #include // Centralized SIMD tier detection + Lane abstractions used by emp-tool's @@ -101,6 +100,8 @@ // // AesLane lives in emp:: (public). ClmulLane lives in emp::detail. +#include "emp-tool/runtime/core/block.h" + #ifdef __x86_64__ namespace emp { diff --git a/emp-tool/runtime/core/transpose.hpp b/emp-tool/runtime/core/transpose.hpp index 0d75c35..79665d1 100644 --- a/emp-tool/runtime/core/transpose.hpp +++ b/emp-tool/runtime/core/transpose.hpp @@ -78,10 +78,6 @@ inline void sse_trans(uint8_t *out, uint8_t const *inp, uint64_t nrows, uint64_t ncols) { uint64_t rr, cc; int i, h; - union { - __m128i x; - uint8_t b[16]; - } tmp; __m128i vec; expecting(nrows % 8 == 0 && ncols % 8 == 0, "sse_trans: dimensions must be multiples of 8"); @@ -131,45 +127,34 @@ inline void sse_trans(uint8_t *out, uint8_t const *inp, uint64_t nrows, if (rr == nrows) return; - // Remainder: 8x(16n+8) bits (n may be 0), processed as pairs of 8x8. - // The non-multiple-of-16 branch uses a scalar variant because the - // 16-wide _mm_set_epi16 pattern below requires aligned 16-row strips. - if ((ncols % 8 == 0 && ncols % 16 != 0) || - (nrows % 8 == 0 && nrows % 16 != 0)) { - for (cc = 0; cc + 16 <= ncols; cc += 16) { - for (i = 0; i < 8; ++i) { - tmp.b[i] = h = detail::load_u16_unaligned(&INP(rr + i, cc)); - tmp.b[i + 8] = h >> 8; - } - for (i = 8; --i >= 0; tmp.x = _mm_slli_epi64(tmp.x, 1)) { - OUT(rr, cc + i) = h = _mm_movemask_epi8(tmp.x); - OUT(rr, cc + i + 8) = h >> 8; - } - } - } else { - for (cc = 0; cc + 16 <= ncols; cc += 16) { - vec = _mm_set_epi16(detail::load_u16_unaligned(&INP(rr + 7, cc)), - detail::load_u16_unaligned(&INP(rr + 6, cc)), - detail::load_u16_unaligned(&INP(rr + 5, cc)), - detail::load_u16_unaligned(&INP(rr + 4, cc)), - detail::load_u16_unaligned(&INP(rr + 3, cc)), - detail::load_u16_unaligned(&INP(rr + 2, cc)), - detail::load_u16_unaligned(&INP(rr + 1, cc)), - detail::load_u16_unaligned(&INP(rr + 0, cc))); - for (i = 8; --i >= 0; vec = _mm_slli_epi64(vec, 1)) { - OUT(rr, cc + i) = h = _mm_movemask_epi8(vec); - OUT(rr, cc + i + 8) = h >> 8; - } + // Exactly 8 rows remain. Process pairs of 8x8 blocks. + for (cc = 0; cc + 16 <= ncols; cc += 16) { + vec = _mm_set_epi16(detail::load_u16_unaligned(&INP(rr + 7, cc)), + detail::load_u16_unaligned(&INP(rr + 6, cc)), + detail::load_u16_unaligned(&INP(rr + 5, cc)), + detail::load_u16_unaligned(&INP(rr + 4, cc)), + detail::load_u16_unaligned(&INP(rr + 3, cc)), + detail::load_u16_unaligned(&INP(rr + 2, cc)), + detail::load_u16_unaligned(&INP(rr + 1, cc)), + detail::load_u16_unaligned(&INP(rr + 0, cc))); + vec = _mm_packus_epi16(_mm_and_si128(vec, _mm_set1_epi16(0xff)), + _mm_srli_epi16(vec, 8)); + for (i = 8; --i >= 0; vec = _mm_slli_epi64(vec, 1)) { + OUT(rr, cc + i) = h = _mm_movemask_epi8(vec); + OUT(rr, cc + i + 8) = h >> 8; } } if (cc == ncols) return; // Do the remaining 8x8 block: - for (i = 0; i < 8; ++i) - tmp.b[i] = INP(rr + i, cc); - for (i = 8; --i >= 0; tmp.x = _mm_slli_epi64(tmp.x, 1)) - OUT(rr, cc + i) = _mm_movemask_epi8(tmp.x); + vec = _mm_set_epi8(0, 0, 0, 0, 0, 0, 0, 0, + INP(rr + 7, cc), INP(rr + 6, cc), + INP(rr + 5, cc), INP(rr + 4, cc), + INP(rr + 3, cc), INP(rr + 2, cc), + INP(rr + 1, cc), INP(rr + 0, cc)); + for (i = 8; --i >= 0; vec = _mm_slli_epi64(vec, 1)) + OUT(rr, cc + i) = _mm_movemask_epi8(vec); } #undef INP #undef OUT diff --git a/emp-tool/runtime/core/utils.h b/emp-tool/runtime/core/utils.h index 613da03..6dec3b0 100644 --- a/emp-tool/runtime/core/utils.h +++ b/emp-tool/runtime/core/utils.h @@ -6,6 +6,7 @@ #include "emp-tool/runtime/core/simd_tier.h" #include #include //https://gcc.gnu.org/gcc-4.9/porting_to.html +#include #include #include // std::_Exit (fatal abort without running destructors) #include @@ -18,6 +19,7 @@ #define macro_str(a) #a namespace emp { + using std::chrono::time_point; using std::chrono::high_resolution_clock; @@ -26,8 +28,8 @@ inline int peer_port(); // $EMP_PORT, default 12345 inline const char * peer_ip(); // $EMP_PEER_IP, default 127.0.0.1 // Timing related -inline time_point clock_start(); -inline double time_from(const time_point& s); +inline std::chrono::steady_clock::time_point clock_start(); +inline double time_from(const std::chrono::steady_clock::time_point& s); // --- Bool / bit packing ------------------------------------------------- diff --git a/emp-tool/runtime/core/utils.hpp b/emp-tool/runtime/core/utils.hpp index ce0ae28..b761bf2 100644 --- a/emp-tool/runtime/core/utils.hpp +++ b/emp-tool/runtime/core/utils.hpp @@ -6,12 +6,14 @@ namespace emp { -inline time_point clock_start() { - return high_resolution_clock::now(); +inline std::chrono::steady_clock::time_point clock_start() { + return std::chrono::steady_clock::now(); } -inline double time_from(const time_point& s) { - return std::chrono::duration_cast(high_resolution_clock::now() - s).count(); +inline double time_from(const std::chrono::steady_clock::time_point& s) { + return std::chrono::duration_cast( + std::chrono::steady_clock::now() - s) + .count(); } // Wait for every future in `res`, then clear it — a barrier folding a batch of @@ -25,7 +27,8 @@ inline void joinNclean(std::vector>& res) { // joinNclean that OR-reduces the bool results (e.g. "did any task flag a cheat?"). inline bool joinNcleanCheat(std::vector>& res) { bool cheat = false; - for (auto& v : res) cheat = cheat || v.get(); + for (auto& v : res) + if (v.get()) cheat = true; res.clear(); return cheat; } @@ -34,16 +37,34 @@ inline bool joinNcleanCheat(std::vector>& res) { // both parties, read from the environment so a two-machine run sets EMP_PORT / // EMP_PEER_IP once per host with no source change. One consequence: two runs on // the same host share EMP_PORT, so don't launch them concurrently. +namespace detail { +inline int parse_bounded_int(const char *text, int lower, int upper, + const char *message) { + expecting(text != nullptr && text[0] != '\0', message); + int value = 0; + const char *end = text + std::strlen(text); + auto parsed = std::from_chars(text, end, value); + expecting(parsed.ec == std::errc{} && parsed.ptr == end && + value >= lower && value <= upper, + message); + return value; +} +} // namespace detail + inline int parse_party(const char *const * arg, int max_party) { - const int p = arg[1] ? atoi(arg[1]) : 0; - expecting(p >= ALICE && p <= max_party, - "parse_party: argv[1] (party) is out of range [1, max_party] " - "(default max is BOB=2; multi-party callers pass nP)"); - return p; + const char *text = arg ? arg[1] : nullptr; + return detail::parse_bounded_int( + text, ALICE, max_party, + "parse_party: argv[1] (party) must be an integer in [1, max_party] " + "(default max is BOB=2; multi-party callers pass nP)"); } inline int peer_port() { const char * e = std::getenv("EMP_PORT"); - return (e && e[0]) ? atoi(e) : 12345; + return (e && e[0]) + ? detail::parse_bounded_int( + e, 1, 65535, + "peer_port: EMP_PORT must be an integer in [1, 65535]") + : 12345; } inline const char * peer_ip() { const char * e = std::getenv("EMP_PEER_IP"); @@ -159,6 +180,8 @@ template inline T bool_to_int(const bool *data) { static_assert(std::is_integral::value, "bool_to_int requires an integral type T"); + static_assert(!std::is_same::type, bool>::value, + "bool_to_int does not support bool"); T ret = 0; bools_to_bits(&ret, data, sizeof(T) * 8); return ret; diff --git a/emp-tool/third_party/ThreadPool.h b/emp-tool/third_party/ThreadPool.h index 881a030..33a7a6d 100644 --- a/emp-tool/third_party/ThreadPool.h +++ b/emp-tool/third_party/ThreadPool.h @@ -32,7 +32,8 @@ freely, subject to the following restrictions: // worker executes the task (see emp-tool/runtime/core/test_mode.h); // (2) enqueue-on-stopped-pool reports through emp::error() instead of // throwing, keeping the public surface exception-free -// (docs/api_conventions.md, enforced by test_no_exceptions). +// (docs/api_conventions.md, enforced by test_no_exceptions); (3) tasks use a +// C++20 capture instead of std::bind, allowing move-only arguments. #include "emp-tool/runtime/core/error.h" #include "emp-tool/runtime/core/test_mode.h" @@ -45,6 +46,7 @@ freely, subject to the following restrictions: #include #include #include +#include #include class ThreadPool { @@ -98,7 +100,14 @@ auto ThreadPool::enqueue(F &&f, Args &&...args) using return_type = typename std::invoke_result::type; auto task = std::make_shared>( - std::bind(std::forward(f), std::forward(args)...)); + [fn = std::forward(f), + ...bound_args = std::forward(args)]() mutable -> return_type { + if constexpr (std::is_invocable_v) + return std::invoke(fn, bound_args...); + else + return std::invoke(std::move(fn), std::move(bound_args)...); + }); std::future res = task->get_future(); // In test mode, the task's lane is derived HERE, on the enqueuing diff --git a/test/runtime/test_block.cpp b/test/runtime/test_block.cpp index 947af1d..2ffbe0f 100644 --- a/test/runtime/test_block.cpp +++ b/test/runtime/test_block.cpp @@ -15,6 +15,7 @@ // bytes_to_bits32 / bits32_to_bytes 32 bools <-> 32 bits // bools_to_bits / bits_to_bools N bools <-> N bits +#include "emp-tool/runtime/core/simd_tier.h" #include "emp-tool/emp-tool.h" #include @@ -33,6 +34,22 @@ using namespace emp; using namespace std; using clk = chrono::high_resolution_clock; +static bool check_direct_simd_tier_include() { +#if EMP_HAS_AVX2 + (void)&emp::detail::sse_trans_n128_avx2; +#endif +#if EMP_HAS_AVX512BW + (void)&emp::detail::sse_trans_n128_avx512bw; +#endif +#if EMP_HAS_GFNI256 && !EMP_HAS_GFNI512 + (void)&emp::detail::sse_trans_n128_gfni256; +#endif +#if EMP_HAS_GFNI512 + (void)&emp::detail::sse_trans_n128_gfni; +#endif + return true; +} + template static bool dies(F&& f) { pid_t pid = fork(); @@ -238,6 +255,35 @@ static bool check_sse_trans_roundtrip() { return true; } +static bool check_sse_trans_scalar_reference() { + PRG prg; + for (uint64_t nrows = 8; nrows <= 128; nrows += 8) { + for (uint64_t ncols = 8; ncols <= 256; ncols += 8) { + const size_t nbytes = (size_t)(nrows * ncols / 8); + const size_t in_stride = (size_t)(ncols / 8); + const size_t out_stride = (size_t)(nrows / 8); + vector in(nbytes), got(nbytes, 0xa5), want(nbytes, 0); + prg.random_data_unaligned(in.data(), (int64_t)nbytes); + + for (uint64_t row = 0; row < nrows; ++row) { + for (uint64_t col = 0; col < ncols; ++col) { + const uint8_t bit = + (in[(size_t)row * in_stride + col / 8] >> (col % 8)) & 1; + want[(size_t)col * out_stride + row / 8] |= bit << (row % 8); + } + } + + sse_trans(got.data(), in.data(), nrows, ncols); + if (got != want) { + cout << " sse_trans scalar-reference FAIL at " + << nrows << "x" << ncols << "\n"; + return false; + } + } + } + return true; +} + // Parity: the tier-dispatched sse_trans_n128 must match the generic // (SSE2-only) sse_trans byte-for-byte. Exercises every variant the build // instantiates (SSE2 / AVX2 / AVX-512BW): on x86 the dispatcher picks the @@ -342,6 +388,7 @@ static bool run_correctness() { cout << "=== correctness ===\n"; struct Case { const char *name; bool (*fn)(); }; Case cases[] = { + {"direct SIMD-tier include", check_direct_simd_tier_include}, {"makeBlock + getLSB", check_makeBlock_getLSB}, {"set_bit", check_set_bit}, {"sigma linear + formula", check_sigma_linear_and_formula}, @@ -350,6 +397,7 @@ static bool run_correctness() { {"cmpBlock", check_cmpBlock}, {"invalid ranges rejected", check_invalid_ranges_rejected}, {"sse_trans round-trip", check_sse_trans_roundtrip}, + {"sse_trans scalar reference", check_sse_trans_scalar_reference}, {"sse_trans_n128 parity", check_sse_trans_n128_parity}, {"bytes<->bits32 round-trip", check_bits_bytes_roundtrip}, {"bool/byte-bools<->bits parity", check_bools_bits_roundtrip}, diff --git a/test/runtime/test_thread_pool.cpp b/test/runtime/test_thread_pool.cpp index 87a6137..d413cea 100644 --- a/test/runtime/test_thread_pool.cpp +++ b/test/runtime/test_thread_pool.cpp @@ -11,6 +11,7 @@ #include #include #include +#include #include #include #include @@ -57,14 +58,42 @@ static bool check_destructor_drains_queue() { return completed.load() == 32; } +static bool check_move_only_argument() { + ThreadPool pool(1); + auto value = make_unique(41); + auto result = pool.enqueue( + [](unique_ptr input) { return *input + 1; }, std::move(value)); + return value == nullptr && result.get() == 42; +} + +static int increment(int& value) { return ++value; } + +struct LvalueCallable { + int operator()() & { return 17; } +}; + +static bool check_legacy_invocation() { + ThreadPool pool(1); + int value = 1; + auto by_reference = pool.enqueue(increment, value); + LvalueCallable callable; + auto lvalue_callable = pool.enqueue(callable); + return by_reference.get() == 2 && value == 1 && + lvalue_callable.get() == 17; +} + static bool run_correctness() { bool tasks = check_tasks_complete(); bool drains = check_destructor_drains_queue(); + bool move_only = check_move_only_argument(); + bool legacy = check_legacy_invocation(); bool rejects_zero = dies([] { ThreadPool pool(0); }); cout << " tasks complete and size is fixed " << (tasks ? "OK" : "FAIL") << "\n"; cout << " destructor drains queued work " << (drains ? "OK" : "FAIL") << "\n"; + cout << " move-only arguments are supported " << (move_only ? "OK" : "FAIL") << "\n"; + cout << " legacy invocation is preserved " << (legacy ? "OK" : "FAIL") << "\n"; cout << " zero workers are rejected " << (rejects_zero ? "OK" : "FAIL") << "\n"; - return tasks && drains && rejects_zero; + return tasks && drains && move_only && legacy && rejects_zero; } int main() { diff --git a/test/runtime/test_utils.cpp b/test/runtime/test_utils.cpp index 59d7aa5..1df6ce1 100644 --- a/test/runtime/test_utils.cpp +++ b/test/runtime/test_utils.cpp @@ -2,6 +2,10 @@ // first; the rest is verification. // // What's in utils.h/utils.hpp (the parts worth testing): +// clock_start(), time_from(start) monotonic elapsed-time helpers +// joinNcleanCheat(futures) wait for all futures and OR their results +// parse_party(argv, max_party) strictly parse a bounded party number +// peer_port() strictly parse EMP_PORT or use 12345 // bool_to_int(const bool *) pack 8*sizeof(T) bools (LSB-first) into T // bool_to_block(const bool *) pack 128 bools into a block @@ -14,11 +18,27 @@ #include #include #include +#include +#include +#include #include using namespace emp; using namespace std; -using clk = chrono::high_resolution_clock; + +template +static bool dies(F&& f) { + pid_t pid = fork(); + if (pid == 0) { + close(STDERR_FILENO); + f(); + _exit(0); + } + if (pid < 0) return false; + int status = 0; + if (waitpid(pid, &status, 0) != pid) return false; + return !WIFEXITED(status) || WEXITSTATUS(status) != 0; +} static void example() { cout << "=== example ===\n"; @@ -42,6 +62,64 @@ static void example() { // ---------- correctness ---------- +static bool check_elapsed_clock() { + static_assert(is_same_v); + auto start = clock_start(); + return time_from(start) >= 0.0; +} + +static bool check_joinNcleanCheat_drains() { + int completed = 0; + vector> results; + results.push_back(async(launch::deferred, [] { return true; })); + results.push_back(async(launch::deferred, [&completed] { + ++completed; + return false; + })); + return joinNcleanCheat(results) && completed == 1 && results.empty(); +} + +static bool check_parse_party() { + const char *alice[] = {"test_utils", "1", nullptr}; + const char *bob[] = {"test_utils", "2", nullptr}; + const char *third[] = {"test_utils", "3", nullptr}; + if (parse_party(alice) != ALICE || parse_party(bob) != BOB || + parse_party(third, 3) != 3) + return false; + + const char *trailing[] = {"test_utils", "2x", nullptr}; + const char *spaced[] = {"test_utils", " 2", nullptr}; + const char *overflow[] = {"test_utils", "999999999999999999999", nullptr}; + const char *out_of_range[] = {"test_utils", "3", nullptr}; + const char *missing[] = {"test_utils", nullptr}; + return dies([&] { parse_party(trailing); }) && + dies([&] { parse_party(spaced); }) && + dies([&] { parse_party(overflow); }) && + dies([&] { parse_party(out_of_range); }) && + dies([&] { parse_party(missing); }); +} + +static bool check_peer_port() { + unsetenv("EMP_PORT"); + if (peer_port() != 12345) return false; + setenv("EMP_PORT", "1", 1); + if (peer_port() != 1) return false; + setenv("EMP_PORT", "65535", 1); + if (peer_port() != 65535) return false; + setenv("EMP_PORT", "", 1); + if (peer_port() != 12345) return false; + + const char *invalid[] = {"0", "65536", "12x", " 12345", + "999999999999999999999"}; + for (const char *value : invalid) { + setenv("EMP_PORT", value, 1); + if (!dies([] { (void)peer_port(); })) return false; + } + unsetenv("EMP_PORT"); + return true; +} + template static bool check_bool_to_int_random(int trials) { PRG prg; @@ -89,20 +167,27 @@ static bool check_bool_to_block_random(int trials) { static bool run_correctness() { cout << "=== correctness ===\n"; - struct Case { const char *name; bool (*fn)(); }; + bool clock = check_elapsed_clock(); + bool joins = check_joinNcleanCheat_drains(); + bool party = check_parse_party(); + bool port = check_peer_port(); auto u8 = []{ return check_bool_to_int_random(64); }; auto u16 = []{ return check_bool_to_int_random(64); }; auto u32 = []{ return check_bool_to_int_random(64); }; auto u64 = []{ return check_bool_to_int_random(64); }; auto blk = []{ return check_bool_to_block_random(64); }; auto kn = []{ return check_bool_to_int_known(); }; + cout << " elapsed clock is monotonic " << (clock ? "OK" : "FAIL") << "\n"; + cout << " joinNcleanCheat drains every future " << (joins ? "OK" : "FAIL") << "\n"; + cout << " parse_party rejects malformed input " << (party ? "OK" : "FAIL") << "\n"; + cout << " peer_port validates EMP_PORT " << (port ? "OK" : "FAIL") << "\n"; bool a = u8(); cout << " bool_to_int random " << (a ? "OK" : "FAIL") << "\n"; bool b = u16(); cout << " bool_to_int random " << (b ? "OK" : "FAIL") << "\n"; bool c = u32(); cout << " bool_to_int random " << (c ? "OK" : "FAIL") << "\n"; bool d = u64(); cout << " bool_to_int random " << (d ? "OK" : "FAIL") << "\n"; bool e = blk(); cout << " bool_to_block random " << (e ? "OK" : "FAIL") << "\n"; bool f = kn(); cout << " bool_to_int<*> known answers " << (f ? "OK" : "FAIL") << "\n"; - return a && b && c && d && e && f; + return clock && joins && party && port && a && b && c && d && e && f; } int main(int /*argc*/, char ** /*argv*/) { From f32f40c67993c71ab31669bbc777330d11d2b9e6 Mon Sep 17 00:00:00 2001 From: Xiao Wang Date: Sat, 22 Aug 2026 08:25:01 -0500 Subject: [PATCH 7/7] Drain task batches and harden ThreadPool lifetime Drain every future before rethrowing the first task failure, preserve exception-disabled builds, and remove unused generic macros from the public umbrella. Make ThreadPool construction exception-safe, derive result types from the stored invocation category, and retain move-only task support. Add direct BlockVec ownership and alignment coverage. --- emp-tool/runtime/core/utils.h | 3 +- emp-tool/runtime/core/utils.hpp | 26 +++++++++ emp-tool/third_party/ThreadPool.h | 59 +++++++++++++++---- test/CMakeLists.txt | 1 + test/runtime/test_block_vector.cpp | 91 +++++++++++++++++++++++++++++ test/runtime/test_no_exceptions.cpp | 7 ++- test/runtime/test_thread_pool.cpp | 27 ++++++++- test/runtime/test_utils.cpp | 51 +++++++++++++++- 8 files changed, 248 insertions(+), 17 deletions(-) create mode 100644 test/runtime/test_block_vector.cpp diff --git a/emp-tool/runtime/core/utils.h b/emp-tool/runtime/core/utils.h index 6dec3b0..0da1bc8 100644 --- a/emp-tool/runtime/core/utils.h +++ b/emp-tool/runtime/core/utils.h @@ -12,11 +12,10 @@ #include #include "emp-tool/runtime/core/constants.h" #include +#include #include #include // joinNclean / joinNcleanCheat: fold a batch of ThreadPool tasks back together #include -#define macro_xstr(a) macro_str(a) -#define macro_str(a) #a namespace emp { diff --git a/emp-tool/runtime/core/utils.hpp b/emp-tool/runtime/core/utils.hpp index b761bf2..5e3295e 100644 --- a/emp-tool/runtime/core/utils.hpp +++ b/emp-tool/runtime/core/utils.hpp @@ -21,15 +21,41 @@ inline double time_from(const std::chrono::steady_clock::time_point& s) { // future alike. template inline void joinNclean(std::vector>& res) { +#if defined(__cpp_exceptions) || defined(__EXCEPTIONS) || defined(_CPPUNWIND) + std::exception_ptr failure; + for (auto& v : res) { + try { + v.get(); + } catch (...) { + if (!failure) failure = std::current_exception(); + } + } + res.clear(); + if (failure) std::rethrow_exception(failure); +#else for (auto& v : res) v.get(); res.clear(); +#endif } // joinNclean that OR-reduces the bool results (e.g. "did any task flag a cheat?"). inline bool joinNcleanCheat(std::vector>& res) { bool cheat = false; +#if defined(__cpp_exceptions) || defined(__EXCEPTIONS) || defined(_CPPUNWIND) + std::exception_ptr failure; + for (auto& v : res) { + try { + if (v.get()) cheat = true; + } catch (...) { + if (!failure) failure = std::current_exception(); + } + } + res.clear(); + if (failure) std::rethrow_exception(failure); +#else for (auto& v : res) if (v.get()) cheat = true; res.clear(); +#endif return cheat; } diff --git a/emp-tool/third_party/ThreadPool.h b/emp-tool/third_party/ThreadPool.h index 33a7a6d..ab9a334 100644 --- a/emp-tool/third_party/ThreadPool.h +++ b/emp-tool/third_party/ThreadPool.h @@ -32,8 +32,8 @@ freely, subject to the following restrictions: // worker executes the task (see emp-tool/runtime/core/test_mode.h); // (2) enqueue-on-stopped-pool reports through emp::error() instead of // throwing, keeping the public surface exception-free -// (docs/api_conventions.md, enforced by test_no_exceptions); (3) tasks use a -// C++20 capture instead of std::bind, allowing move-only arguments. +// (docs/api_conventions.md, enforced by test_no_exceptions); (3) queued tasks +// own their callable and arguments and may therefore contain move-only values. #include "emp-tool/runtime/core/error.h" #include "emp-tool/runtime/core/test_mode.h" @@ -49,16 +49,46 @@ freely, subject to the following restrictions: #include #include +namespace emp::detail { + +template +struct thread_pool_stored_result; + +template +struct thread_pool_stored_result + : std::invoke_result &, std::decay_t &...> {}; + +template +struct thread_pool_stored_result + : std::invoke_result &&, std::decay_t &&...> {}; + +template +using thread_pool_stored_result_t = typename thread_pool_stored_result< + std::is_invocable_v &, std::decay_t &...>, F, + Args...>::type; + +} // namespace emp::detail + class ThreadPool { public: ThreadPool(size_t); template auto enqueue(F &&f, Args &&...args) - -> std::future::type>; + -> std::future>; ~ThreadPool(); size_t size() const; private: + struct construction_guard { + ThreadPool *pool; + ~construction_guard() { + if (pool != nullptr) pool->stop_and_join(); + } + void release() noexcept { pool = nullptr; } + }; + + void stop_and_join() noexcept; + // need to keep track of threads so we can join them std::vector workers; // the task queue @@ -75,6 +105,8 @@ inline size_t ThreadPool::size() const { return workers.size(); } // the constructor just launches some amount of workers inline ThreadPool::ThreadPool(size_t threads) : stop(false) { emp::expecting(threads > 0, "ThreadPool: worker count must be positive"); + workers.reserve(threads); + construction_guard guard{this}; for (size_t i = 0; i < threads; ++i) workers.emplace_back([this] { for (;;) { @@ -91,22 +123,23 @@ inline ThreadPool::ThreadPool(size_t threads) : stop(false) { task(); } }); + guard.release(); } // add new work item to the pool template auto ThreadPool::enqueue(F &&f, Args &&...args) - -> std::future::type> { - using return_type = typename std::invoke_result::type; + -> std::future> { + using return_type = emp::detail::thread_pool_stored_result_t; auto task = std::make_shared>( [fn = std::forward(f), ...bound_args = std::forward(args)]() mutable -> return_type { - if constexpr (std::is_invocable_v) - return std::invoke(fn, bound_args...); - else - return std::invoke(std::move(fn), std::move(bound_args)...); + if constexpr (std::is_invocable_v) + return std::invoke(fn, bound_args...); + else + return std::invoke(std::move(fn), std::move(bound_args)...); }); std::future res = task->get_future(); @@ -132,8 +165,7 @@ auto ThreadPool::enqueue(F &&f, Args &&...args) return res; } -// the destructor joins all threads -inline ThreadPool::~ThreadPool() { +inline void ThreadPool::stop_and_join() noexcept { { std::unique_lock lock(queue_mutex); stop = true; @@ -142,4 +174,7 @@ inline ThreadPool::~ThreadPool() { for (std::thread &worker : workers) worker.join(); } +// the destructor joins all threads +inline ThreadPool::~ThreadPool() { stop_and_join(); } + #endif diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index 653b51e..a44fba0 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -62,6 +62,7 @@ endif() # ---- runtime/ : core, crypto, io, execution primitives ---- add_test_case(runtime test_aes) add_test_case(runtime test_block) +add_test_case(runtime test_block_vector) add_test_case(runtime test_error) add_test_case(runtime test_ccrh) add_test_case(runtime test_f2k) diff --git a/test/runtime/test_block_vector.cpp b/test/runtime/test_block_vector.cpp new file mode 100644 index 0000000..8b93625 --- /dev/null +++ b/test/runtime/test_block_vector.cpp @@ -0,0 +1,91 @@ +// core/block_vector.h — aligned block storage without redundant zero-fill. +// Read example() first; the rest is verification. +// +// What's in block_vector.h: +// default_init_allocator default-initialize storage for overwrite-first use +// default_init_vector vector using that allocator +// BlockVec aligned, overwrite-first block vector + +#include "emp-tool/runtime/core/block_vector.h" + +#include +#include +#include +#include + +using namespace emp; +using namespace std; + +static bool blocks_equal(block lhs, block rhs) { + return cmpBlock(&lhs, &rhs, 1); +} + +static void example() { + BlockVec blocks(2); + blocks[0] = makeBlock(0, 1); + blocks[1] = makeBlock(0, 2); + cout << "BlockVec size = " << blocks.size() + << ", first = " << blocks[0] << "\n"; +} + +static bool check_alignment() { + for (size_t size : {size_t{1}, size_t{3}, size_t{257}}) { + BlockVec blocks(size); + if (reinterpret_cast(blocks.data()) % alignof(block) != 0) + return false; + } + return true; +} + +static bool check_copy_move_reserve_resize() { + BlockVec original(4); + for (size_t i = 0; i < original.size(); ++i) + original[i] = makeBlock(i + 10, i + 1); + + BlockVec copy = original; + BlockVec moved = std::move(copy); + moved.reserve(32); + for (size_t i = 0; i < original.size(); ++i) + if (!blocks_equal(moved[i], original[i])) return false; + + moved.resize(8); + for (size_t i = 4; i < moved.size(); ++i) + moved[i] = makeBlock(i + 10, i + 1); + for (size_t i = 0; i < moved.size(); ++i) + if (!blocks_equal(moved[i], makeBlock(i + 10, i + 1))) return false; + + moved.resize(2); + return moved.size() == 2 && + blocks_equal(moved[0], makeBlock(10, 1)) && + blocks_equal(moved[1], makeBlock(11, 2)); +} + +struct NonTrivial { + NonTrivial() : value(23) { ++constructions; } + static inline int constructions = 0; + int value; +}; + +static bool check_nontrivial_default_construction() { + NonTrivial::constructions = 0; + default_init_vector values(3); + return NonTrivial::constructions == 3 && values[0].value == 23 && + values[1].value == 23 && values[2].value == 23; +} + +static bool run_correctness() { + static_assert(is_same_v); + bool alignment = check_alignment(); + bool lifetime = check_copy_move_reserve_resize(); + bool nontrivial = check_nontrivial_default_construction(); + cout << " block alignment is preserved " << (alignment ? "OK" : "FAIL") << "\n"; + cout << " copy/move/reserve/resize are sound " << (lifetime ? "OK" : "FAIL") << "\n"; + cout << " nontrivial default ctor runs " << (nontrivial ? "OK" : "FAIL") << "\n"; + return alignment && lifetime && nontrivial; +} + +int main() { + example(); + cout << "\n=== correctness ===\n"; + return run_correctness() ? 0 : 1; +} diff --git a/test/runtime/test_no_exceptions.cpp b/test/runtime/test_no_exceptions.cpp index 2afdb44..0c80fcf 100644 --- a/test/runtime/test_no_exceptions.cpp +++ b/test/runtime/test_no_exceptions.cpp @@ -4,7 +4,8 @@ // from the public umbrella fails the BUILD, not just the run. The body // instantiates the two paths that historically threw: ThreadPool::enqueue // (vendored, patched to emp::error) and PRG system-entropy seeding -// (std::random_device, replaced by getentropy). +// (std::random_device, replaced by getentropy). It also instantiates the +// exception-disabled future-draining helpers. #include "emp-tool/emp-tool.h" #include using namespace emp; @@ -20,6 +21,10 @@ int main() { std::printf("test_no_exceptions: pool result mismatch\n"); return 1; } + std::vector> empty; + joinNclean(empty); + std::vector> empty_cheat; + if (joinNcleanCheat(empty_cheat)) return 1; std::printf("test_no_exceptions: OK\n"); return 0; } diff --git a/test/runtime/test_thread_pool.cpp b/test/runtime/test_thread_pool.cpp index d413cea..c7c6e89 100644 --- a/test/runtime/test_thread_pool.cpp +++ b/test/runtime/test_thread_pool.cpp @@ -12,7 +12,9 @@ #include #include #include +#include #include +#include #include #include @@ -72,6 +74,20 @@ struct LvalueCallable { int operator()() & { return 17; } }; +struct CategoryOverload { + int operator()(int& value) { return value + 1; } + string operator()(int&&) { return "wrong overload"; } +}; + +struct NotCallable {}; + +template +concept PoolEnqueueable = requires(ThreadPool& pool, F&& fn) { + pool.enqueue(std::forward(fn)); +}; + +static_assert(!PoolEnqueueable); + static bool check_legacy_invocation() { ThreadPool pool(1); int value = 1; @@ -82,18 +98,27 @@ static bool check_legacy_invocation() { lvalue_callable.get() == 17; } +static bool check_result_matches_stored_invocation() { + ThreadPool pool(1); + auto result = pool.enqueue(CategoryOverload{}, 41); + static_assert(is_same_v>); + return result.get() == 42; +} + static bool run_correctness() { bool tasks = check_tasks_complete(); bool drains = check_destructor_drains_queue(); bool move_only = check_move_only_argument(); bool legacy = check_legacy_invocation(); + bool stored_result = check_result_matches_stored_invocation(); bool rejects_zero = dies([] { ThreadPool pool(0); }); cout << " tasks complete and size is fixed " << (tasks ? "OK" : "FAIL") << "\n"; cout << " destructor drains queued work " << (drains ? "OK" : "FAIL") << "\n"; cout << " move-only arguments are supported " << (move_only ? "OK" : "FAIL") << "\n"; cout << " legacy invocation is preserved " << (legacy ? "OK" : "FAIL") << "\n"; + cout << " stored invocation type is exact " << (stored_result ? "OK" : "FAIL") << "\n"; cout << " zero workers are rejected " << (rejects_zero ? "OK" : "FAIL") << "\n"; - return tasks && drains && move_only && legacy && rejects_zero; + return tasks && drains && move_only && legacy && stored_result && rejects_zero; } int main() { diff --git a/test/runtime/test_utils.cpp b/test/runtime/test_utils.cpp index 1df6ce1..c3d4bcd 100644 --- a/test/runtime/test_utils.cpp +++ b/test/runtime/test_utils.cpp @@ -3,13 +3,21 @@ // // What's in utils.h/utils.hpp (the parts worth testing): // clock_start(), time_from(start) monotonic elapsed-time helpers +// joinNclean(futures) wait for every future, including after failure // joinNcleanCheat(futures) wait for all futures and OR their results +// umbrella inclusion preserve caller-owned generic macro names // parse_party(argv, max_party) strictly parse a bounded party number // peer_port() strictly parse EMP_PORT or use 12345 // bool_to_int(const bool *) pack 8*sizeof(T) bools (LSB-first) into T // bool_to_block(const bool *) pack 128 bools into a block +#define macro_str(value) 17 +#define macro_xstr(value) 19 #include "emp-tool/emp-tool.h" +static_assert(macro_str(user_owned) == 17); +static_assert(macro_xstr(user_owned) == 19); +#undef macro_str +#undef macro_xstr #include #include @@ -17,6 +25,7 @@ #include #include #include +#include #include #include #include @@ -80,6 +89,41 @@ static bool check_joinNcleanCheat_drains() { return joinNcleanCheat(results) && completed == 1 && results.empty(); } +static bool check_joinNclean_exception_drains() { + int completed = 0; + vector> results; + results.push_back(async(launch::deferred, [] { + throw runtime_error("first failure"); + })); + results.push_back(async(launch::deferred, [&completed] { ++completed; })); + try { + joinNclean(results); + } catch (const runtime_error& e) { + return string(e.what()) == "first failure" && completed == 1 && + results.empty(); + } + return false; +} + +static bool check_joinNcleanCheat_exception_drains() { + int completed = 0; + vector> results; + results.push_back(async(launch::deferred, []() -> bool { + throw runtime_error("cheat failure"); + })); + results.push_back(async(launch::deferred, [&completed] { + ++completed; + return true; + })); + try { + (void)joinNcleanCheat(results); + } catch (const runtime_error& e) { + return string(e.what()) == "cheat failure" && completed == 1 && + results.empty(); + } + return false; +} + static bool check_parse_party() { const char *alice[] = {"test_utils", "1", nullptr}; const char *bob[] = {"test_utils", "2", nullptr}; @@ -169,6 +213,8 @@ static bool run_correctness() { cout << "=== correctness ===\n"; bool clock = check_elapsed_clock(); bool joins = check_joinNcleanCheat_drains(); + bool join_failure = check_joinNclean_exception_drains(); + bool cheat_failure = check_joinNcleanCheat_exception_drains(); bool party = check_parse_party(); bool port = check_peer_port(); auto u8 = []{ return check_bool_to_int_random(64); }; @@ -179,6 +225,8 @@ static bool run_correctness() { auto kn = []{ return check_bool_to_int_known(); }; cout << " elapsed clock is monotonic " << (clock ? "OK" : "FAIL") << "\n"; cout << " joinNcleanCheat drains every future " << (joins ? "OK" : "FAIL") << "\n"; + cout << " joinNclean drains after an exception " << (join_failure ? "OK" : "FAIL") << "\n"; + cout << " cheat join drains after an exception " << (cheat_failure ? "OK" : "FAIL") << "\n"; cout << " parse_party rejects malformed input " << (party ? "OK" : "FAIL") << "\n"; cout << " peer_port validates EMP_PORT " << (port ? "OK" : "FAIL") << "\n"; bool a = u8(); cout << " bool_to_int random " << (a ? "OK" : "FAIL") << "\n"; @@ -187,7 +235,8 @@ static bool run_correctness() { bool d = u64(); cout << " bool_to_int random " << (d ? "OK" : "FAIL") << "\n"; bool e = blk(); cout << " bool_to_block random " << (e ? "OK" : "FAIL") << "\n"; bool f = kn(); cout << " bool_to_int<*> known answers " << (f ? "OK" : "FAIL") << "\n"; - return clock && joins && party && port && a && b && c && d && e && f; + return clock && joins && join_failure && cheat_failure && party && port && + a && b && c && d && e && f; } int main(int /*argc*/, char ** /*argv*/) {