From cece21c15f5d315a7518d1c4e2c85e5065a46123 Mon Sep 17 00:00:00 2001 From: wzu Date: Sat, 26 Sep 2026 08:56:47 +0800 Subject: [PATCH 1/7] client-rs: multi-rail RDMA read path for a single Worker Add MultiRailClient: N local RDMA devices (rails) read one striped object in parallel over existing stripe-subset descriptor GETs. - Rail management: per-rail verbs ctx/PD/CQ/QP/cached-MR, health + cooldown recovery, endpoint whitelists (rail-optimized pinning), sysfs topology (NUMA/PCIe) reporting. - Planning: placement validation (missing/dup/out-of-bounds stripes before any I/O), byte-exact weighted stripe allocation (cross- multiplied least-loaded), wave dispatch under connection caps. - Safety: control-plane io/connect timeouts; failed reads quiesce connections (Stop -> join -> QP destroy -> MR dereg) before returning, so late server WRITEs cannot touch freed/reused memory; safe API evicts cached registrations synchronously. - Integrity: per-task byte/chunk-count checks against the layout, StaleDescriptor typed error on generation/etag mismatch, optional per-stripe xxh3-64 verification matching the server encoding. - Backpressure: per-rail/total connection and in-flight-byte limits. - Observability: RailSnapshot per-rail stats. - Tools/tests: cs-multirail-bench, 21 hardware-independent unit tests, hardware-gated e2e (dual-rail correctness, failure injection + buffer-reuse safety, stale descriptor). - rdma.rs: additive io/connect timeouts, device listing, GID query, num_chunks GET outcome, synchronous MR eviction; fix upstream rdma-bench test compile + clippy nits; allow result_large_err on tonic-generated stubs (newer clippy). --- Cargo.lock | 1 + kv-service/client-rs/Cargo.toml | 8 +- .../client-rs/src/bin/multirail_bench.rs | 328 +++ kv-service/client-rs/src/bin/rdma_bench.rs | 2 +- kv-service/client-rs/src/lib.rs | 9 + kv-service/client-rs/src/multirail.rs | 1981 +++++++++++++++++ kv-service/client-rs/src/rdma.rs | 158 +- kv-service/client-rs/tests/multirail_e2e.rs | 301 +++ kv-service/server/src/lib.rs | 5 + 9 files changed, 2785 insertions(+), 8 deletions(-) create mode 100644 kv-service/client-rs/src/bin/multirail_bench.rs create mode 100644 kv-service/client-rs/src/multirail.rs create mode 100644 kv-service/client-rs/tests/multirail_e2e.rs diff --git a/Cargo.lock b/Cargo.lock index 5caefbb..d063aef 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -355,6 +355,7 @@ dependencies = [ "tokio-stream", "tonic", "tonic-build", + "twox-hash", ] [[package]] diff --git a/kv-service/client-rs/Cargo.toml b/kv-service/client-rs/Cargo.toml index 8ad21a0..8a4350c 100644 --- a/kv-service/client-rs/Cargo.toml +++ b/kv-service/client-rs/Cargo.toml @@ -14,6 +14,11 @@ name = "cs-rdma-bench" path = "src/bin/rdma_bench.rs" required-features = ["rdma"] +[[bin]] +name = "cs-multirail-bench" +path = "src/bin/multirail_bench.rs" +required-features = ["rdma"] + [dependencies] # gRPC (versions aligned with the server) tonic = "0.11" @@ -29,10 +34,11 @@ clap = { version = "4", features = ["derive"] } rdma-sys = { version = "0.3", optional = true } anyhow = { version = "1", optional = true } libc = { version = "0.2", optional = true } +twox-hash = { version = "1.6", optional = true } [features] default = [] -rdma = ["dep:rdma-sys", "dep:anyhow", "dep:libc"] +rdma = ["dep:rdma-sys", "dep:anyhow", "dep:libc", "dep:twox-hash"] [build-dependencies] tonic-build = "0.11" diff --git a/kv-service/client-rs/src/bin/multirail_bench.rs b/kv-service/client-rs/src/bin/multirail_bench.rs new file mode 100644 index 0000000..45d595c --- /dev/null +++ b/kv-service/client-rs/src/bin/multirail_bench.rs @@ -0,0 +1,328 @@ +//! Multi-rail benchmark / functional driver for the ContextStore client. +//! +//! Reads one striped object through N local RDMA devices in parallel and +//! reports per-rail throughput, error, and in-flight metrics, plus the +//! aggregate bandwidth and (optionally) a single-rail baseline for the +//! multi-rail speedup ratio. +//! +//! Example (two Soft-RoCE rails against a two-NIC server): +//! ```text +//! cs-multirail-bench --coordinator http://10.0.0.1:50051 \ +//! --namespace bench --object-key obj1 --put-size-mb 1024 --verify \ +//! --rails "rxe0,rxe1" --iters 5 --baseline +//! ``` + +use anyhow::{anyhow, Context, Result}; +use clap::Parser; +use contextstore_client_rs::multirail::{ + MultiRailClient, RailConfig, RailLimits, RailSelectPolicy, +}; +use contextstore_client_rs::{KvClient, ObjectLookup}; +use prost::bytes::Bytes; +use std::alloc::{alloc_zeroed, dealloc, Layout}; +use std::collections::HashMap; +use std::time::Instant; + +#[derive(Parser, Debug)] +#[command(about = "Multi-rail RDMA read benchmark")] +struct Args { + /// gRPC coordinator endpoint, e.g. http://10.0.0.1:50051 + #[arg(long)] + coordinator: String, + #[arg(long)] + namespace: String, + #[arg(long)] + object_key: String, + /// Comma-separated rail specs: device[:port[:gid[:weight]]] + #[arg(long, value_delimiter = ',')] + rails: Vec, + /// Explicit rail↔endpoint pinning, e.g. `rxe0=10.0.0.2:50054,rxe1=10.0.0.1:50053` + /// (rail-optimized fabrics / cross-wired testbeds). Multiple endpoints per + /// rail separated by `;`. + #[arg(long, value_delimiter = ',')] + pin: Vec, + /// Rewrite per-stripe RDMA endpoints round-robin across this list (single + /// node exposing several NIC listeners; every listener serves every stripe + /// of the node). Enables true multi-rail reads against one storage node. + #[arg(long, value_delimiter = ',')] + alternate_endpoints: Vec, + /// Stripe→rail policy: least-loaded | endpoint-affinity | rr + #[arg(long, default_value = "least-loaded")] + policy: String, + #[arg(long, default_value = "5")] + iters: usize, + /// Destination buffer size in MiB (>= object size). + #[arg(long, default_value = "1024")] + buf_mb: usize, + /// Seed the object with a deterministic pattern first (needed for --verify). + #[arg(long)] + put_size_mb: Option, + /// Verify every byte of every read against the deterministic pattern. + #[arg(long, default_value_t = true)] + verify: bool, + /// Also run a single-rail baseline for the speedup ratio. + #[arg(long, default_value_t = true)] + baseline: bool, + #[arg(long, default_value = "30")] + io_timeout_secs: u64, +} + +/// Deterministic object pattern: each 8-byte word derived from its index. +fn pattern_word(word_index: u64) -> u64 { + word_index + .wrapping_mul(0x9E37_79B9_7F4A_7C15) + .rotate_left(23) + ^ word_index +} + +// chunks_exact keeps a single implementation shared with the verifier below; +// as_chunks_mut's tuple API would obscure it. +#[allow(clippy::chunks_exact_to_as_chunks)] +fn fill_pattern(buffer: &mut [u8]) { + for (index, chunk) in buffer.chunks_exact_mut(8).enumerate() { + chunk.copy_from_slice(&pattern_word(index as u64).to_le_bytes()); + } + let words = buffer.len() / 8 * 8; + for (i, byte) in buffer[words..].iter_mut().enumerate() { + *byte = pattern_word(words as u64 / 8 + i as u64) as u8; + } +} + +#[allow(clippy::chunks_exact_to_as_chunks)] +fn verify_pattern(buffer: &[u8], label: &str) -> Result<()> { + for (index, chunk) in buffer.chunks_exact(8).enumerate() { + let expected = pattern_word(index as u64).to_le_bytes(); + if chunk != expected { + let first_bad = index * 8; + return Err(anyhow!( + "{label}: verification failed at byte {first_bad} (0x{:x?} != 0x{:x?})", + &chunk[..4], + &expected[..4] + )); + } + } + Ok(()) +} + +struct AlignedBuffer { + ptr: *mut u8, + layout: Layout, + len: usize, +} + +// The pointer is only dereferenced while the struct is alive; the multi-rail +// read keeps it registered and quiesced within that window. +unsafe impl Send for AlignedBuffer {} + +impl AlignedBuffer { + fn new(len: usize) -> Result { + let layout = Layout::from_size_align(len, 4096)?; + let ptr = unsafe { alloc_zeroed(layout) }; + if ptr.is_null() { + return Err(anyhow!("failed to allocate {len} byte buffer")); + } + Ok(Self { ptr, layout, len }) + } + + fn as_mut(&mut self) -> &mut [u8] { + unsafe { std::slice::from_raw_parts_mut(self.ptr, self.len) } + } +} + +impl Drop for AlignedBuffer { + fn drop(&mut self) { + unsafe { dealloc(self.ptr, self.layout) }; + } +} + +fn seed_object(args: &Args, runtime: &tokio::runtime::Runtime, size_mb: usize) -> Result<()> { + let size = size_mb * 1024 * 1024; + let mut buffer = AlignedBuffer::new(size)?; + fill_pattern(buffer.as_mut()); + runtime.block_on(async { + let mut client = KvClient::connect(format_coordinator(&args.coordinator)) + .await + .map_err(|error| anyhow!(error.to_string()))?; + let big = Bytes::from(buffer.as_mut().to_vec()); + let chunk = 4 * 1024 * 1024; + let mut segments = Vec::new(); + for offset in (0..size).step_by(chunk) { + segments.push(big.slice(offset..(offset + chunk).min(size))); + } + client + .put_stream_chunks(&args.namespace, &args.object_key, segments) + .await + .map_err(|error| anyhow!(error.to_string()))?; + Ok::<_, anyhow::Error>(()) + })?; + println!("[seed] wrote {size} byte object via gRPC"); + Ok(()) +} + +fn format_coordinator(url: &str) -> String { + if url.starts_with("http://") || url.starts_with("https://") { + url.to_string() + } else { + format!("http://{url}") + } +} + +fn lookup(args: &Args, runtime: &tokio::runtime::Runtime) -> Result { + runtime.block_on(async { + let mut client = KvClient::connect(format_coordinator(&args.coordinator)) + .await + .map_err(|error| anyhow!(error.to_string()))?; + client + .lookup_object(&args.namespace, &args.object_key) + .await + .map_err(|error| anyhow!(error.to_string())) + })? + .ok_or_else(|| anyhow!("object not found: {}/{}", args.namespace, args.object_key)) +} + +fn run_client( + args: &Args, + lookup: &ObjectLookup, + rails: Vec, + label: &str, +) -> Result<(f64, usize)> { + let policy = RailSelectPolicy::parse(&args.policy) + .with_context(|| format!("unknown policy '{}'", args.policy))?; + let limits = RailLimits { + io_timeout: std::time::Duration::from_secs(args.io_timeout_secs), + ..RailLimits::default() + }; + let client = MultiRailClient::new(rails)?.with_limits(limits).with_policy(policy); + let mut buffer = AlignedBuffer::new(args.buf_mb * 1024 * 1024)?; + let object_size = usize::try_from(lookup.descriptor.size)?; + + println!( + "[{label}] rails={} policy={:?} object={}B stripes={} chunk={}B", + client.rail_count(), + policy, + object_size, + lookup.descriptor.stripe_count, + lookup.descriptor.chunk_size, + ); + for snapshot in client.rails_snapshot() { + println!("[{label}] {snapshot}"); + } + + let mut latencies = Vec::with_capacity(args.iters); + let mut last_len = 0usize; + for iteration in 0..args.iters { + // Poison the buffer so a missing stripe cannot slip through. + buffer.as_mut().iter_mut().for_each(|b| *b = 0xA5); + let started = Instant::now(); + let bytes = client + .read_lookup_into(lookup, buffer.as_mut()) + .with_context(|| format!("[{label}] iteration {iteration} failed"))?; + latencies.push(started.elapsed()); + last_len = bytes; + if bytes != object_size { + return Err(anyhow!( + "[{label}] iteration {iteration}: got {bytes} bytes, expected {object_size}" + )); + } + if args.verify { + verify_pattern(&buffer.as_mut()[..object_size], &format!("{label}#{iteration}"))?; + } + } + + for snapshot in client.rails_snapshot() { + println!("[{label}] {snapshot}"); + } + let total: f64 = latencies.iter().map(|d| d.as_secs_f64()).sum(); + let best = latencies.iter().map(|d| d.as_secs_f64()).fold(f64::INFINITY, f64::min); + let gbps = object_size as f64 / 1024f64.powi(3) / (total / latencies.len() as f64); + println!( + "[{label}] avg={:.3}s best={:.3}s avg_bw={:.3} GiB/s over {} iterations", + total / latencies.len() as f64, + best, + gbps, + latencies.len() + ); + Ok((gbps, last_len)) +} + +fn parse_pins(specs: &[String]) -> HashMap> { + let mut pins = HashMap::new(); + for spec in specs { + if let Some((device, targets)) = spec.split_once('=') { + pins.insert( + device.trim().to_string(), + targets + .split(';') + .map(|t| t.trim().to_string()) + .filter(|t| !t.is_empty()) + .collect(), + ); + } + } + pins +} + +fn main() -> Result<()> { + let args = Args::parse(); + if args.rails.is_empty() { + return Err(anyhow!("--rails is required, e.g. --rails rxe0,rxe1")); + } + let runtime = tokio::runtime::Runtime::new()?; + + if let Some(size_mb) = args.put_size_mb { + seed_object(&args, &runtime, size_mb)?; + } + let mut lookup = lookup(&args, &runtime)?; + if lookup.placement.is_none() { + return Err(anyhow!("lookup returned no placement (is the object striped?)")); + } + let original_lookup = lookup.clone(); + if !args.alternate_endpoints.is_empty() { + // Spread the node's stripes over its listeners so several rails can + // carry the object concurrently. + let endpoints = &args.alternate_endpoints; + if let Some(placement) = lookup.placement.as_mut() { + for (index, chunk) in placement.chunks.iter_mut().enumerate() { + chunk.rdma_endpoint = endpoints[index % endpoints.len()].clone(); + } + } + } + + let pins = parse_pins(&args.pin); + let rails: Vec = args + .rails + .iter() + .map(|spec| RailConfig::parse(spec).ok_or_else(|| anyhow!("bad rail spec '{spec}'"))) + .map(|rail| { + rail.map(|mut rail| { + if let Some(endpoints) = pins.get(&rail.device) { + rail.endpoints = endpoints.clone(); + } + rail + }) + }) + .collect::>()?; + + // Multi-rail run. + let (multi_gbps, bytes) = run_client(&args, &lookup, rails.clone(), "multi")?; + + // Optional single-rail baseline for the speedup ratio (unpinned, reading + // the original placement: one rail must reach the advertised endpoint). + if args.baseline && rails.len() > 1 { + let mut baseline_rail = rails[0].clone(); + baseline_rail.endpoints.clear(); + let (single_gbps, single_bytes) = + run_client(&args, &original_lookup, vec![baseline_rail], "single")?; + if single_bytes != bytes { + return Err(anyhow!("single/multi rail byte counts differ")); + } + println!( + "[speedup] single={:.3} GiB/s multi({})={:.3} GiB/s ratio={:.2}x", + single_gbps, + rails.len(), + multi_gbps, + multi_gbps / single_gbps + ); + } + Ok(()) +} diff --git a/kv-service/client-rs/src/bin/rdma_bench.rs b/kv-service/client-rs/src/bin/rdma_bench.rs index 75cc4c6..85ac6ce 100644 --- a/kv-service/client-rs/src/bin/rdma_bench.rs +++ b/kv-service/client-rs/src/bin/rdma_bench.rs @@ -694,7 +694,7 @@ fn run_multi_endpoint(args: &Args, coordinator: &str) -> Result<()> { // 把整个目标 buffer 均分为 N 段 (同一 MR, 不同偏移), // 走 tag-15 SGE 路径; server 按段映射逐段 WRITE. let view = registered.view(); - let seg_len = (buf_size / sge_segments).max(1) as u64; + let seg_len = (buf_size / sge_segments.max(1)).max(1) as u64; let mut segments = Vec::with_capacity(sge_segments); let mut off = 0u64; let (base, rkey, total) = diff --git a/kv-service/client-rs/src/lib.rs b/kv-service/client-rs/src/lib.rs index 096266c..ef0d675 100644 --- a/kv-service/client-rs/src/lib.rs +++ b/kv-service/client-rs/src/lib.rs @@ -8,6 +8,11 @@ //! concatenating 480MB on the client side. Instead, the gRPC framework's inbound //! buffer view is handed straight to the caller. +// tonic-build's generated client stubs return Result<_, tonic::Status>; newer +// clippy versions flag the Status size on that generated code. The generated +// file is vendored from the proto, so silence the lint crate-wide. +#![allow(clippy::result_large_err)] + pub mod pb { tonic::include_proto!("contextstore.kv.v1"); } @@ -19,6 +24,10 @@ pub mod pb { #[cfg(feature = "rdma")] pub mod rdma; +/// Multi-rail RDMA read path: parallel stripe reads over several local HCAs. +#[cfg(feature = "rdma")] +pub mod multirail; + use pb::kv_service_client::KvServiceClient; use prost::bytes::Bytes; use tonic::transport::Channel; diff --git a/kv-service/client-rs/src/multirail.rs b/kv-service/client-rs/src/multirail.rs new file mode 100644 index 0000000..8d37df3 --- /dev/null +++ b/kv-service/client-rs/src/multirail.rs @@ -0,0 +1,1981 @@ +//! Multi-rail RDMA read path: one client Worker reads a single striped object +//! in parallel over several local HCAs ("rails"). +//! +//! # Model +//! +//! A *rail* is one local RDMA device (HCA port) plus its slice of per-rail +//! resources: verbs context, PD, CQ, QPs and cached memory registrations. The +//! remote side is addressed by the per-stripe RDMA endpoints carried in a +//! `PlacementDescriptor` — the server already exposes one listener per NIC +//! (`CS_RDMA_DEVICES`), so a rail↔endpoint pairing is simply a connection. +//! +//! Reads keep the existing data path untouched: every connection sends one +//! stripe-subset descriptor GET over its TCP control stream and the server +//! RDMA-WRITEs the requested stripes at `base + stripe_index * chunk_size` of +//! the client-registered destination. Because stripes land in disjoint +//! regions, rails and endpoints transfer concurrently with no client-side +//! reassembly. +//! +//! The destination buffer is registered **once per rail device** (an rkey is +//! only meaningful to the device it was registered on), so a read through N +//! rails holds N MRs of the same memory. +//! +//! # Memory-safety contract (late-write protection) +//! +//! RDMA WRITEs on the GET path are server-initiated; a failed or timed-out +//! control exchange does not prove the WRITE never happened. This module +//! guarantees that [`MultiRailClient::read_object_into`] (and the `_raw` +//! variant) returns only when no work request targeting the caller's buffer +//! can still place data: +//! +//! * success: every task received its tag-3 response, which the server only +//! sends after its last RDMA WRITE completed — the RC transport has already +//! placed all bytes; +//! * failure/timeout: the affected connections are quiesced — `Stop` is +//! queued behind any in-flight command, the worker thread is joined (so the +//! QP is destroyed and the cached MRs are deregistered, in that order), +//! and only then does the call return. Stray retransmissions hit a +//! destroyed QP / invalid rkey and are dropped by the transport instead of +//! DMA-ing into freed or reused memory. +//! +//! This is the "safe failure, no in-request transparent retry" posture: the +//! first version fails the whole read when a rail breaks; the rail enters a +//! cooldown window and later reads recover by planning around it. +//! +//! # Backpressure +//! +//! [`RailLimits`] bounds total connections, per-rail connections (queue +//! depth), and in-flight bytes (per rail and total). Planning groups tasks +//! into waves so a wave never exceeds the connection caps, and dispatch +//! blocks — with a deadline — until byte headroom exists. +//! +//! # Integrity +//! +//! Every task carries the full object descriptor (handle / generation / ETag +//! / layout version); the server rejects mismatches by answering +//! `found=false`, which surfaces here as [`MultiRailError::StaleDescriptor`]. +//! After the bytes arrive the client re-derives the expected per-stripe byte +//! count from the placement, cross-checks the server-reported `num_chunks`, +//! and, when the placement carries per-stripe xxh3-64 checksums, verifies +//! every stripe it received. Missing, duplicated, or out-of-bounds stripes +//! are detected while planning, before any network I/O is issued. + +use crate::pb; +use crate::rdma::{self, RdmaClient, RdmaClientConfig}; +use std::collections::{BTreeMap, HashMap, HashSet}; +use std::fmt; +use std::path::Path; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::{mpsc, Arc, Mutex}; +use std::thread::JoinHandle; +use std::time::{Duration, Instant}; + +/// Cooldown applied to a rail after a failure before it is considered for +/// planning again (path recovery window). +const DEFAULT_RAIL_COOLDOWN: Duration = Duration::from_secs(10); + +// --------------------------------------------------------------------------- +// Configuration +// --------------------------------------------------------------------------- + +/// One local rail: an RDMA device + port + GID to use for it. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct RailConfig { + /// Verbs device name, e.g. `mlx5_0` or `rxe0`. + pub device: String, + /// HCA port number (almost always 1). + pub port: u8, + /// GID index for the address handle. + pub gid_index: u8, + /// Relative share of stripes this rail receives (load-balancing weight). + pub weight: u32, + /// Optional endpoint whitelist (`host:port` or bare `host`). Empty means + /// the rail may serve any endpoint. Rail-optimized fabrics pin rail k to + /// fabric k; soft-RoCE cross-wired testbeds use it the same way. + pub endpoints: Vec, +} + +impl RailConfig { + pub fn new(device: impl Into) -> Self { + Self { + device: device.into(), + port: 1, + gid_index: 3, + weight: 1, + endpoints: Vec::new(), + } + } + + pub fn with_port(mut self, port: u8) -> Self { + self.port = port; + self + } + + pub fn with_gid_index(mut self, gid_index: u8) -> Self { + self.gid_index = gid_index; + self + } + + pub fn with_weight(mut self, weight: u32) -> Self { + self.weight = weight.max(1); + self + } + + /// Restrict this rail to the given endpoints (`host:port` or bare host). + pub fn with_endpoints(mut self, endpoints: Vec) -> Self { + self.endpoints = endpoints; + self + } + + /// Whether this rail may serve `endpoint` (exact `host:port` or host-only + /// match; an empty whitelist allows everything). + fn allows_endpoint(&self, endpoint: &str) -> bool { + if self.endpoints.is_empty() { + return true; + } + let host = endpoint.rsplit_once(':').map(|(h, _)| h).unwrap_or(endpoint); + self.endpoints + .iter() + .any(|allowed| allowed == endpoint || *allowed == host) + } + + /// Parse `device[:port[:gid[:weight]]]` (portions left out keep defaults). + pub fn parse(spec: &str) -> Option { + let mut parts = spec.split(':'); + let device = parts.next()?.trim(); + if device.is_empty() { + return None; + } + let mut config = Self::new(device); + if let Some(port) = parts.next().and_then(|p| p.trim().parse().ok()) { + config = config.with_port(port); + } + if let Some(gid) = parts.next().and_then(|g| g.trim().parse().ok()) { + config = config.with_gid_index(gid); + } + if let Some(weight) = parts.next().and_then(|w| w.trim().parse().ok()) { + config = config.with_weight(weight); + } + Some(config) + } +} + +/// Resource and backpressure limits for multi-rail reads. +#[derive(Clone, Debug)] +pub struct RailLimits { + /// Maximum simultaneous connections (queue depth) one rail may hold. + pub max_connections_per_rail: usize, + /// Maximum simultaneous connections across all rails. + pub max_connections_total: usize, + /// In-flight bytes one rail may accumulate before dispatch blocks. + pub max_inflight_bytes_per_rail: u64, + /// In-flight bytes across all rails before dispatch blocks. + pub max_inflight_bytes_total: u64, + /// Per-operation TCP control-channel timeout; also bounds connect. + pub io_timeout: Duration, + /// How long a failed rail is skipped before it is retried. + pub rail_cooldown: Duration, +} + +impl Default for RailLimits { + fn default() -> Self { + Self { + max_connections_per_rail: 8, + max_connections_total: 32, + max_inflight_bytes_per_rail: 8 * 1024 * 1024 * 1024u64, + max_inflight_bytes_total: 32 * 1024 * 1024 * 1024u64, + io_timeout: Duration::from_secs(30), + rail_cooldown: DEFAULT_RAIL_COOLDOWN, + } + } +} + +/// Stripe→rail assignment strategy. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum RailSelectPolicy { + /// Spread every endpoint's stripes over all healthy rails, weighted by + /// rail weight and current assignment load (default; maximum fan-out). + #[default] + LeastLoaded, + /// Pin each endpoint to a single rail — prefer a subnet match between the + /// rail GID and the endpoint IP, then the least-loaded rail. Fewer + /// connections, cleaner topology, no fan-out within one endpoint. + EndpointAffinity, + /// Pure weighted round-robin over healthy rails. + WeightedRoundRobin, +} + +impl RailSelectPolicy { + pub fn parse(value: &str) -> Option { + match value.trim().to_ascii_lowercase().as_str() { + "least-loaded" | "least_loaded" | "leastloaded" => Some(Self::LeastLoaded), + "endpoint-affinity" | "endpoint_affinity" | "affinity" => Some(Self::EndpointAffinity), + "weighted-round-robin" | "round-robin" | "rr" => Some(Self::WeightedRoundRobin), + _ => None, + } + } +} + +// --------------------------------------------------------------------------- +// Errors +// --------------------------------------------------------------------------- + +/// Typed failures of a multi-rail read. The `StaleDescriptor` variant means +/// the server rejected the descriptor identity (generation / ETag / layout +/// changed) and the caller should re-run the gRPC lookup before retrying. +#[derive(Debug)] +pub enum MultiRailError { + NoRails, + NoHealthyRails, + BufferTooSmall { need: u64, have: usize }, + InvalidPlacement(String), + MissingStripes(Vec), + DuplicateStripes(Vec), + OutOfBoundsStripe { stripe: u32, stripe_count: u32 }, + StaleDescriptor { endpoint: String }, + TaskFailed { + rail: String, + endpoint: String, + source: String, + }, + Timeout { + rail: String, + endpoint: String, + after_ms: u128, + }, + ByteCountMismatch { + rail: String, + endpoint: String, + expected: u64, + actual: u64, + }, + ChunkCountMismatch { + rail: String, + endpoint: String, + expected: u32, + actual: u32, + }, + ChecksumMismatch { + stripe: u32, + expected: String, + actual: String, + }, + WorkerPanic { rail: String, endpoint: String }, +} + +impl fmt::Display for MultiRailError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::NoRails => write!(f, "multi-rail client has no rails configured"), + Self::NoHealthyRails => write!(f, "all rails are unhealthy (in cooldown)"), + Self::BufferTooSmall { need, have } => write!( + f, + "destination buffer too small: need {need} bytes, have {have}" + ), + Self::InvalidPlacement(reason) => write!(f, "invalid placement: {reason}"), + Self::MissingStripes(stripes) => { + write!(f, "placement is missing stripes {stripes:?}") + } + Self::DuplicateStripes(stripes) => { + write!(f, "placement duplicates stripes {stripes:?}") + } + Self::OutOfBoundsStripe { stripe, stripe_count } => write!( + f, + "stripe index {stripe} out of bounds for stripe count {stripe_count}" + ), + Self::StaleDescriptor { endpoint } => write!( + f, + "server at {endpoint} rejected the descriptor (stale generation/etag/layout); re-lookup required" + ), + Self::TaskFailed { + rail, + endpoint, + source, + } => write!(f, "rail {rail} failed reading from {endpoint}: {source}"), + Self::Timeout { + rail, + endpoint, + after_ms, + } => write!(f, "rail {rail} timed out reading from {endpoint} after {after_ms}ms"), + Self::ByteCountMismatch { + rail, + endpoint, + expected, + actual, + } => write!( + f, + "rail {rail} got {actual} bytes from {endpoint}, expected {expected}" + ), + Self::ChunkCountMismatch { + rail, + endpoint, + expected, + actual, + } => write!( + f, + "rail {rail}: server at {endpoint} served {actual} chunks, expected {expected}" + ), + Self::ChecksumMismatch { + stripe, + expected, + actual, + } => write!( + f, + "checksum mismatch on stripe {stripe}: expected {expected}, got {actual}" + ), + Self::WorkerPanic { rail, endpoint } => { + write!(f, "rail {rail} worker for {endpoint} panicked") + } + } + } +} + +impl std::error::Error for MultiRailError {} + +// --------------------------------------------------------------------------- +// Topology +// --------------------------------------------------------------------------- + +/// Host topology attributes of a rail, read from sysfs. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct RailTopology { + /// NUMA node the device is attached to (-1 when unknown). + pub numa_node: i32, + /// PCIe BDF, e.g. `0000:60:00.0` (empty when unknown). + pub pci_slot: String, +} + +impl fmt::Display for RailTopology { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + f, + "numa={}", + if self.numa_node < 0 { + "?".to_string() + } else { + self.numa_node.to_string() + } + )?; + if !self.pci_slot.is_empty() { + write!(f, " pci={}", self.pci_slot)?; + } + Ok(()) + } +} + +/// Read a rail's topology from `//device`. +fn topology_from(base: &Path, device: &str) -> RailTopology { + let dev_dir = base.join(device); + let numa = std::fs::read_to_string(dev_dir.join("device/numa_node")) + .ok() + .and_then(|text| text.trim().parse::().ok()) + .unwrap_or(-1); + let pci_slot = dev_dir + .canonicalize() + .ok() + .and_then(|path| { + path.file_name() + .map(|name| name.to_string_lossy().into_owned()) + }) + .unwrap_or_default(); + RailTopology { numa_node: numa, pci_slot } +} + +/// Read a rail's topology from the real sysfs tree. +pub fn read_topology(device: &str) -> RailTopology { + topology_from(Path::new("/sys/class/infiniband"), device) +} + +// --------------------------------------------------------------------------- +// Per-rail stats & state +// --------------------------------------------------------------------------- + +#[derive(Default)] +struct RailStats { + requests_ok: AtomicU64, + requests_err: AtomicU64, + bytes_read: AtomicU64, + timeouts: AtomicU64, + connections_created: AtomicU64, + connections_quiesced: AtomicU64, + inflight_requests: AtomicU64, + inflight_bytes: AtomicU64, +} + +/// Point-in-time view of one rail for observability. +#[derive(Clone, Debug)] +pub struct RailSnapshot { + pub index: usize, + pub device: String, + pub topology: RailTopology, + pub healthy: bool, + pub cooldown_ms_remaining: u64, + pub requests_ok: u64, + pub requests_err: u64, + pub bytes_read: u64, + pub timeouts: u64, + pub connections_created: u64, + pub connections_quiesced: u64, + pub inflight_requests: u64, + pub inflight_bytes: u64, +} + +impl fmt::Display for RailSnapshot { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + f, + "rail[{}] {} {} healthy={} cooldown={}ms ok={} err={} tmo={} read={}MiB conns=+{}/~{} inflight=(req:{},B:{})", + self.index, + self.device, + self.topology, + self.healthy, + self.cooldown_ms_remaining, + self.requests_ok, + self.requests_err, + self.timeouts, + self.bytes_read / (1024 * 1024), + self.connections_created, + self.connections_quiesced, + self.inflight_requests, + self.inflight_bytes, + ) + } +} + +struct Rail { + index: usize, + config: RailConfig, + topology: RailTopology, + stats: RailStats, + healthy: AtomicBool, + unhealthy_until: Mutex>, + /// IPv4 address embedded in the rail GID, when the GID is an IPv4-mapped + /// RoCE v2 GID (used for endpoint affinity). + gid_v4: Option, +} + +impl Rail { + fn mark_unhealthy(&self, cooldown: Duration) { + self.healthy.store(false, Ordering::Release); + *self.unhealthy_until.lock().unwrap() = Some(Instant::now() + cooldown); + } + + fn cooldown_remaining(&self) -> Duration { + self.unhealthy_until + .lock() + .unwrap() + .map(|until| until.saturating_duration_since(Instant::now())) + .unwrap_or_default() + } + + /// A rail is selectable when healthy, or when its cooldown expired. + fn selectable(&self) -> bool { + if self.healthy.load(Ordering::Acquire) { + return true; + } + if self.cooldown_remaining().is_zero() { + // Cooldown over: allow retry and restore health optimistically; + // a new failure puts it straight back into cooldown. + *self.unhealthy_until.lock().unwrap() = None; + self.healthy.store(true, Ordering::Release); + return true; + } + false + } +} + +// --------------------------------------------------------------------------- +// Placement validation & read planning +// --------------------------------------------------------------------------- + +#[derive(Debug)] +struct StripeInfo { + endpoint: Arc, + offset: u64, + length: u64, + checksum: Option, +} + +#[derive(Debug)] +struct ValidatedPlacement { + is_striped: bool, + stripe_count: u32, + object_size: u64, + /// Stripe index → placement record. Covers exactly `0..stripe_count` for + /// striped objects; a single implicit stripe 0 for plain objects. + stripes: BTreeMap, + expected_bytes: u64, +} + +/// Byte length of `stripe` under a striped layout (last stripe is short). +fn stripe_length(stripe: u32, stripe_count: u32, chunk_size: u64, size: u64) -> Option { + if chunk_size == 0 || stripe >= stripe_count { + return None; + } + let base = stripe as u64 * chunk_size; + if base >= size { + return None; + } + Some((size - base).min(chunk_size)) +} + +/// Validate descriptor ↔ placement consistency: every stripe present exactly +/// once, indices in range, offsets/lengths matching the striped layout, and +/// every stripe owned by a non-empty RDMA endpoint. Detects the missing / +/// duplicated / out-of-bounds cases before any network I/O happens. +fn validate_placement( + descriptor: &pb::ObjectDescriptor, + chunks: &[pb::PlacementChunk], +) -> Result { + let object_size = descriptor.size; + if chunks.is_empty() { + return Err(MultiRailError::InvalidPlacement( + "placement contains no chunks".into(), + )); + } + if descriptor.is_striped && (descriptor.chunk_size == 0 || descriptor.stripe_count == 0) { + return Err(MultiRailError::InvalidPlacement( + "striped descriptor with zero chunk_size/stripe_count".into(), + )); + } + + let stripe_count = if descriptor.is_striped { + descriptor.stripe_count + } else { + 1 + }; + let chunk_size = if descriptor.is_striped { + descriptor.chunk_size + } else { + object_size + }; + + let mut stripes: BTreeMap = BTreeMap::new(); + let mut expected_bytes = 0u64; + for chunk in chunks { + let stripe = chunk.stripe_index; + if stripe >= stripe_count { + return Err(MultiRailError::OutOfBoundsStripe { stripe, stripe_count }); + } + if chunk.rdma_endpoint.is_empty() { + return Err(MultiRailError::InvalidPlacement(format!( + "stripe {stripe} has no RDMA endpoint" + ))); + } + let expected_len = if descriptor.is_striped { + stripe_length(stripe, stripe_count, chunk_size, object_size).ok_or( + MultiRailError::OutOfBoundsStripe { stripe, stripe_count }, + )? + } else { + object_size + }; + if chunk.length != 0 && chunk.length != expected_len { + return Err(MultiRailError::InvalidPlacement(format!( + "stripe {stripe} length {} does not match layout expectation {expected_len}", + chunk.length + ))); + } + let expected_offset = stripe as u64 * chunk_size; + if chunk.offset != 0 && chunk.offset != expected_offset { + return Err(MultiRailError::InvalidPlacement(format!( + "stripe {stripe} offset {} does not match natural position {expected_offset}", + chunk.offset + ))); + } + if stripes + .insert( + stripe, + StripeInfo { + endpoint: Arc::from(chunk.rdma_endpoint.as_str()), + offset: expected_offset, + length: expected_len, + checksum: if chunk.checksum.is_empty() { + None + } else { + Some(chunk.checksum.clone()) + }, + }, + ) + .is_some() + { + return Err(MultiRailError::DuplicateStripes(vec![stripe])); + } + expected_bytes += expected_len; + } + + let missing: Vec = (0..stripe_count) + .filter(|index| !stripes.contains_key(index)) + .collect(); + if !missing.is_empty() { + return Err(MultiRailError::MissingStripes(missing)); + } + + Ok(ValidatedPlacement { + is_striped: descriptor.is_striped, + stripe_count, + object_size, + stripes, + expected_bytes, + }) +} + +/// One unit of dispatch: a stripe subset read from one endpoint over one rail. +#[derive(Debug)] +struct TaskSpec { + rail_index: usize, + endpoint: Arc, + stripes: Vec, + bytes: u64, +} + +/// The planned distribution of a read across rails. +#[derive(Debug)] +pub struct ReadPlan { + tasks: Vec, + expected_bytes: u64, + stripe_count: u32, +} + +impl ReadPlan { + pub fn task_count(&self) -> usize { + self.tasks.len() + } + + pub fn expected_bytes(&self) -> u64 { + self.expected_bytes + } + + pub fn stripe_count(&self) -> u32 { + self.stripe_count + } +} + +/// Extract the IPv4 host part of an `host:port` endpoint string. +fn endpoint_v4(endpoint: &str) -> Option { + let host = endpoint + .rsplit_once(':') + .map(|(host, _)| host) + .unwrap_or(endpoint); + host.parse().ok() +} + +/// If the GID bytes are an IPv4-mapped RoCE v2 GID, return the v4 address. +fn gid_to_v4(raw: [u8; 16]) -> Option { + if raw[..10].iter().all(|b| *b == 0) && raw[10] == 0xff && raw[11] == 0xff { + return Some(std::net::Ipv4Addr::new(raw[12], raw[13], raw[14], raw[15])); + } + None +} + +/// Pick the candidate rail with the smallest `assigned / weight` fraction, +/// comparing by cross-multiplication (no integer division rounding) and +/// breaking ties toward the higher-weight rail. +fn least_loaded(candidates: &[usize], rails: &[Arc], assigned: &[u64]) -> usize { + let mut best = candidates[0]; + for &candidate in &candidates[1..] { + // assigned ≤ 2^40 and weight ≤ 2^32, so the products fit u128. + let lhs = (assigned[candidate] as u128) * (rails[best].config.weight as u128); + let rhs = (assigned[best] as u128) * (rails[candidate].config.weight as u128); + if lhs < rhs + || (lhs == rhs && rails[candidate].config.weight > rails[best].config.weight) + { + best = candidate; + } + } + best +} + +/// Choose the rail for one stripe (or one endpoint, under `EndpointAffinity`) +/// out of `candidates` (indices into `rails`, pre-filtered by endpoint +/// whitelists). +fn pick_rail( + endpoint: &str, + rails: &[Arc], + candidates: &[usize], + assigned: &[u64], + policy: RailSelectPolicy, + rr: &mut u64, +) -> usize { + match policy { + RailSelectPolicy::WeightedRoundRobin => { + let total: u32 = candidates.iter().map(|&i| rails[i].config.weight).sum(); + if total == 0 || candidates.is_empty() { + return candidates.first().copied().unwrap_or(0); + } + let mut slot = *rr % total as u64; + *rr += 1; + for &index in candidates { + let weight = rails[index].config.weight as u64; + if slot < weight { + return index; + } + slot -= weight; + } + candidates[0] + } + RailSelectPolicy::EndpointAffinity => { + // Prefer rails on the same /24 as the endpoint, then the least + // loaded relative to weight. + let mut shortlist: Vec = candidates.to_vec(); + if let Some(ip) = endpoint_v4(endpoint) { + let matched: Vec = candidates + .iter() + .copied() + .filter(|&index| { + rails[index].gid_v4 + .is_some_and(|rail_ip| rail_ip.octets()[..3] == ip.octets()[..3]) + }) + .collect(); + if !matched.is_empty() { + shortlist = matched; + } + } + least_loaded(&shortlist, rails, assigned) + } + RailSelectPolicy::LeastLoaded => least_loaded(candidates, rails, assigned), + } +} + +/// Build the task list for a validated placement: group stripes by endpoint, +/// then distribute each endpoint's stripes over the rails allowed to reach +/// that endpoint, according to the policy, keeping one task per (rail, +/// endpoint) pair so the byte count per task can be verified afterwards. +fn build_plan( + placement: &ValidatedPlacement, + rails: &[Arc], + policy: RailSelectPolicy, +) -> Result { + // Group stripes by endpoint first. + let mut by_endpoint: BTreeMap, Vec<(u32, u64)>> = BTreeMap::new(); + for (stripe, info) in &placement.stripes { + by_endpoint + .entry(Arc::clone(&info.endpoint)) + .or_default() + .push((*stripe, info.length)); + } + + let mut assigned = vec![0u64; rails.len()]; + let mut rr = 0u64; + let mut tasks: Vec = Vec::new(); + for (endpoint, mut stripes) in by_endpoint { + let candidates: Vec = rails + .iter() + .enumerate() + .filter(|(_, rail)| rail.config.allows_endpoint(&endpoint)) + .map(|(index, _)| index) + .collect(); + if candidates.is_empty() { + return Err(MultiRailError::InvalidPlacement(format!( + "no rail is allowed to reach endpoint {endpoint} (check rail endpoint whitelists)" + ))); + } + stripes.sort_unstable(); + match policy { + // Non-striped objects are read whole: the server's stripe-subset + // path rejects them, so the task carries an empty stripe list and + // falls back to the full-object descriptor GET. + _ if !placement.is_striped => { + tasks.push(TaskSpec { + rail_index: candidates[0], + endpoint, + stripes: Vec::new(), + bytes: placement.expected_bytes, + }); + } + // Split this endpoint's stripes over every allowed rail. + RailSelectPolicy::LeastLoaded | RailSelectPolicy::WeightedRoundRobin => { + let mut per_rail: Vec> = vec![Vec::new(); rails.len()]; + let mut bytes_per_rail = vec![0u64; rails.len()]; + for (stripe, length) in stripes { + let rail = pick_rail(&endpoint, rails, &candidates, &assigned, policy, &mut rr); + per_rail[rail].push(stripe); + bytes_per_rail[rail] += length; + assigned[rail] += length; + } + for (rail_index, (stripe_list, bytes)) in + per_rail.into_iter().zip(bytes_per_rail).enumerate() + { + if stripe_list.is_empty() { + continue; + } + tasks.push(TaskSpec { + rail_index, + endpoint: Arc::clone(&endpoint), + stripes: stripe_list, + bytes, + }); + } + } + // Pin the whole endpoint to one allowed rail. + RailSelectPolicy::EndpointAffinity => { + let rail = pick_rail(&endpoint, rails, &candidates, &assigned, policy, &mut rr); + let bytes: u64 = stripes.iter().map(|(_, length)| *length).sum(); + assigned[rail] += bytes; + tasks.push(TaskSpec { + rail_index: rail, + endpoint, + stripes: stripes.into_iter().map(|(stripe, _)| stripe).collect(), + bytes, + }); + } + } + } + + Ok(ReadPlan { + tasks, + expected_bytes: placement.expected_bytes, + stripe_count: placement.stripe_count, + }) +} + +/// Group tasks into dispatch waves that respect the connection caps. Tasks +/// never share a (rail, endpoint) pair within one plan, so every task in a +/// wave is a distinct connection. +fn plan_waves(tasks: Vec, limits: &RailLimits) -> Vec> { + let mut waves: Vec> = Vec::new(); + let mut current: Vec = Vec::new(); + let mut rail_conns: HashMap = HashMap::new(); + for task in tasks { + let rail_ok = rail_conns.get(&task.rail_index).copied().unwrap_or(0) + < limits.max_connections_per_rail.max(1); + let total_ok = current.len() < limits.max_connections_total.max(1); + if !rail_ok || !total_ok { + waves.push(std::mem::take(&mut current)); + rail_conns.clear(); + } + *rail_conns.entry(task.rail_index).or_insert(0) += 1; + current.push(task); + } + if !current.is_empty() { + waves.push(current); + } + if waves.is_empty() { + waves.push(Vec::new()); + } + waves +} + +// --------------------------------------------------------------------------- +// Connection workers +// --------------------------------------------------------------------------- + +enum Command { + Read { + descriptor: Arc, + stripes: Arc>, + dst_base: usize, + dst_len: usize, + reply: mpsc::Sender, + }, + EvictRegistration { + base: usize, + ack: mpsc::Sender<()>, + }, + Stop, +} + +/// Result of one dispatched task. The worker computes the *expected* byte and +/// chunk counts from the descriptor + requested stripes so the dispatcher can +/// verify completeness without external bookkeeping. +struct TaskReply { + rail_index: usize, + endpoint: Arc, + expected_bytes: u64, + expected_chunks: u32, + outcome: Result, String>, +} + +impl TaskReply { + /// Best-effort classification of transport errors caused by timeouts. + fn is_timeout(&self) -> bool { + match &self.outcome { + Err(message) => { + message.contains("timed out") + || message.contains("WouldBlock") + || message.contains("etimedout") + || message.contains("ETIMEDOUT") + || message.contains("os error 110") + } + Ok(_) => false, + } + } +} + +struct ConnEntry { + rail_index: usize, + tx: mpsc::Sender, + handle: JoinHandle<()>, +} + +/// Expected bytes/chunks for one stripe-subset request against `descriptor`. +fn expected_task_outcome( + descriptor: &pb::ObjectDescriptor, + stripes: &[u32], +) -> (u64, u32) { + if stripes.is_empty() { + // Non-striped (or whole-object) GET: one implicit chunk. + return (descriptor.size, 1); + } + let mut total = 0u64; + for &stripe in stripes { + let length = if descriptor.is_striped { + stripe_length( + stripe, + descriptor.stripe_count, + descriptor.chunk_size, + descriptor.size, + ) + .unwrap_or(0) + } else { + descriptor.size + }; + total += length; + } + (total, stripes.len() as u32) +} + +fn connect_rail_client( + rail: &Rail, + endpoint: &str, + limits: &RailLimits, +) -> Result { + let config = RdmaClientConfig::new(endpoint.to_string(), rail.config.device.clone()) + .with_port(rail.config.port) + .with_gid_index(rail.config.gid_index) + .with_io_timeout(limits.io_timeout) + .with_connect_timeout(limits.io_timeout); + RdmaClient::connect(config).map_err(|error| error.to_string()) +} + +/// Worker body for one (rail, endpoint) connection. Owns the `RdmaClient` +/// exclusively; commands arrive serialized, so the TCP control channel +/// invariant of one request in flight per connection holds. +/// +/// Teardown ordering is the memory-safety core: `Stop` exits the loop, then +/// `RdmaClient` drops — BYE + `ibv_destroy_qp` first, cached MRs +/// deregistered after (each MR's `Arc` keeps the verbs context +/// alive until the last registration is gone). +fn conn_worker( + rail: Arc, + endpoint: Arc, + limits: RailLimits, + rx: mpsc::Receiver, + setup_tx: mpsc::Sender>, +) { + let mut client = match connect_rail_client(&rail, &endpoint, &limits) { + Ok(client) => { + rail.stats.connections_created.fetch_add(1, Ordering::Relaxed); + let _ = setup_tx.send(Ok(())); + client + } + Err(error) => { + let _ = setup_tx.send(Err(error)); + // Drain until Stop so the dispatcher's join never blocks on a + // command nobody will read. + for command in rx.try_iter() { + if matches!(command, Command::Stop) { + break; + } + } + return; + } + }; + + while let Ok(command) = rx.recv() { + match command { + Command::Read { + descriptor, + stripes, + dst_base, + dst_len, + reply, + } => { + let (expected_bytes, expected_chunks) = + expected_task_outcome(&descriptor, &stripes); + let outcome = (|| -> Result, String> { + // SAFETY: dst_base..dst_base+dst_len stays valid and + // unmoved for the whole read call — the dispatcher joins + // every worker before returning, and registrations for + // non-sticky buffers are evicted synchronously at the end + // of each read. + let view = unsafe { + client + .register_raw_buffer_cached(dst_base as *mut u8, dst_len) + .map_err(|error| error.to_string())? + }; + client + .get_descriptor_stripes_into_view_detailed(&descriptor, &stripes, view, 0) + .map_err(|error| error.to_string()) + })(); + let _ = reply.send(TaskReply { + rail_index: rail.index, + endpoint: Arc::clone(&endpoint), + expected_bytes, + expected_chunks, + outcome, + }); + } + Command::EvictRegistration { base, ack } => { + client.evict_registrations_for(base); + let _ = ack.send(()); + } + Command::Stop => break, + } + } +} + +// --------------------------------------------------------------------------- +// MultiRailClient +// --------------------------------------------------------------------------- + +/// A multi-rail RDMA reader for one client Worker. +/// +/// Holds one connection per (rail, endpoint) pair used so far. Connections +/// persist across reads; a failed connection is quiesced (worker joined, QP +/// destroyed) before the failing read returns, and the rail enters a cooldown +/// window that later reads plan around. +pub struct MultiRailClient { + rails: Vec>, + limits: RailLimits, + policy: RailSelectPolicy, + conns: Mutex), ConnEntry>>, +} + +impl MultiRailClient { + /// Create a client over the given rails. Fails when no rail is given. + pub fn new(rails: Vec) -> Result { + if rails.is_empty() { + return Err(MultiRailError::NoRails); + } + let rails = rails + .into_iter() + .enumerate() + .map(|(index, config)| { + let topology = read_topology(&config.device); + let gid_v4 = + rdma::query_gid_raw(&config.device, config.port, config.gid_index) + .and_then(gid_to_v4); + Arc::new(Rail { + index, + topology, + gid_v4, + config, + stats: RailStats::default(), + healthy: AtomicBool::new(true), + unhealthy_until: Mutex::new(None), + }) + }) + .collect(); + Ok(Self { + rails, + limits: RailLimits::default(), + policy: RailSelectPolicy::default(), + conns: Mutex::new(HashMap::new()), + }) + } + + /// Override resource/backpressure limits. + pub fn with_limits(mut self, limits: RailLimits) -> Self { + self.limits = limits; + self + } + + /// Override the stripe→rail selection policy. + pub fn with_policy(mut self, policy: RailSelectPolicy) -> Self { + self.policy = policy; + self + } + + pub fn rail_count(&self) -> usize { + self.rails.len() + } + + pub fn limits(&self) -> &RailLimits { + &self.limits + } + + /// Snapshot of per-rail state for observability (health, throughput, + /// errors, in-flight counters, topology). + pub fn rails_snapshot(&self) -> Vec { + let mut live_conns: HashMap = HashMap::new(); + for key in self.conns.lock().unwrap().keys() { + *live_conns.entry(key.0).or_insert(0) += 1; + } + self.rails + .iter() + .map(|rail| { + let s = &rail.stats; + RailSnapshot { + index: rail.index, + device: rail.config.device.clone(), + topology: rail.topology.clone(), + healthy: rail.healthy.load(Ordering::Acquire), + cooldown_ms_remaining: rail.cooldown_remaining().as_millis() as u64, + requests_ok: s.requests_ok.load(Ordering::Relaxed), + requests_err: s.requests_err.load(Ordering::Relaxed), + bytes_read: s.bytes_read.load(Ordering::Relaxed), + timeouts: s.timeouts.load(Ordering::Relaxed), + connections_created: s.connections_created.load(Ordering::Relaxed), + connections_quiesced: s.connections_quiesced.load(Ordering::Relaxed), + inflight_requests: s.inflight_requests.load(Ordering::Relaxed), + inflight_bytes: s.inflight_bytes.load(Ordering::Relaxed), + } + }) + .collect() + } + + /// Read an object described by `descriptor` + `placement chunks` into + /// `buffer`, distributing stripes across all healthy rails. + /// + /// Returns the number of bytes placed in `buffer` (== descriptor.size on + /// success). On error the whole read fails safely: every connection that + /// had work in flight is quiesced before this function returns, so no + /// late RDMA WRITE can touch `buffer` after the call. + pub fn read_object_into( + &self, + descriptor: &pb::ObjectDescriptor, + chunks: &[pb::PlacementChunk], + buffer: &mut [u8], + ) -> Result { + let base = buffer.as_mut_ptr() as usize; + self.read_impl(descriptor, chunks, base, buffer.len(), false) + } + + /// Convenience wrapper taking a gRPC [`crate::ObjectLookup`] result. + pub fn read_lookup_into( + &self, + lookup: &crate::ObjectLookup, + buffer: &mut [u8], + ) -> Result { + let placement = lookup.placement.as_ref().ok_or_else(|| { + MultiRailError::InvalidPlacement("lookup returned no placement".into()) + })?; + self.read_object_into(&lookup.descriptor, &placement.chunks, buffer) + } + + /// `read_object_into` for FFI-owned / pinned memory. + /// + /// # Safety + /// `ptr..ptr+len` must be valid writable memory that stays alive, unmoved + /// and unmodified for the duration of this call. With + /// `sticky_registration = true` the per-rail registrations of this buffer + /// are kept cached in the connections (skipping ~1.5 ms `ibv_reg_mr` per + /// rail on subsequent reads); the caller then guarantees the buffer is a + /// long-lived pool region that is never freed nor reused for non-RDMA + /// purposes while this client lives. + pub unsafe fn read_object_into_raw( + &self, + descriptor: &pb::ObjectDescriptor, + chunks: &[pb::PlacementChunk], + ptr: *mut u8, + len: usize, + sticky_registration: bool, + ) -> Result { + if ptr.is_null() || len == 0 { + return Err(MultiRailError::BufferTooSmall { + need: descriptor.size, + have: len, + }); + } + self.read_impl(descriptor, chunks, ptr as usize, len, sticky_registration) + } + + fn read_impl( + &self, + descriptor: &pb::ObjectDescriptor, + chunks: &[pb::PlacementChunk], + base: usize, + len: usize, + sticky: bool, + ) -> Result { + let placement = validate_placement(descriptor, chunks)?; + if placement.object_size > len as u64 { + return Err(MultiRailError::BufferTooSmall { + need: placement.object_size, + have: len, + }); + } + + // Plan around healthy rails only. + let rails: Vec> = self + .rails + .iter() + .filter(|rail| rail.selectable()) + .cloned() + .collect(); + if rails.is_empty() { + return Err(MultiRailError::NoHealthyRails); + } + let rail_of: HashMap> = rails + .iter() + .map(|rail| (rail.index, Arc::clone(rail))) + .collect(); + + let plan = build_plan(&placement, &rails, self.policy)?; + let waves = plan_waves(plan.tasks, &self.limits); + + // Global deadline: every wave gets a full io_timeout, plus one extra + // for connect phases. + let deadline = Instant::now() + + self + .limits + .io_timeout + .mul_f64(waves.len().max(1) as f64) + + self.limits.io_timeout; + + let (reply_tx, reply_rx) = mpsc::channel::(); + let descriptor = Arc::new(descriptor.clone()); + let mut participated: HashSet<(usize, Arc)> = HashSet::new(); + let mut failure: Option = None; + let mut failed_rail: Option = None; + let mut verified_bytes = 0u64; + + 'waves: for wave in waves { + // ---- dispatch this wave ---- + let mut dispatched = 0usize; + for task in &wave { + let Some(rail) = rail_of.get(&task.rail_index) else { + continue; + }; + // Backpressure: wait for per-rail and total byte headroom. + if !self.await_headroom(rail, task.bytes, deadline) { + failure = Some(MultiRailError::Timeout { + rail: rail.config.device.clone(), + endpoint: task.endpoint.to_string(), + after_ms: self.limits.io_timeout.as_millis(), + }); + failed_rail = Some(rail.index); + break 'waves; + } + match self.get_or_create_conn(rail, Arc::clone(&task.endpoint), deadline) { + Ok(tx) => { + participated.insert((rail.index, Arc::clone(&task.endpoint))); + rail.stats.inflight_requests.fetch_add(1, Ordering::Relaxed); + rail.stats.inflight_bytes.fetch_add(task.bytes, Ordering::Relaxed); + let sent = tx.send(Command::Read { + descriptor: Arc::clone(&descriptor), + stripes: Arc::new(task.stripes.clone()), + dst_base: base, + dst_len: len, + reply: reply_tx.clone(), + }); + if sent.is_err() { + rail.stats.inflight_requests.fetch_sub(1, Ordering::Relaxed); + rail.stats.inflight_bytes.fetch_sub(task.bytes, Ordering::Relaxed); + failure = Some(MultiRailError::TaskFailed { + rail: rail.config.device.clone(), + endpoint: task.endpoint.to_string(), + source: "worker exited unexpectedly".into(), + }); + failed_rail = Some(rail.index); + break 'waves; + } + dispatched += 1; + } + Err(error) => { + failed_rail = Some(rail.index); + failure = Some(error); + break 'waves; + } + } + } + + // ---- collect this wave's replies (one per dispatched task) ---- + let mut replies: Vec = Vec::with_capacity(dispatched); + while replies.len() < dispatched { + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + failure = Some(MultiRailError::Timeout { + rail: "any".into(), + endpoint: "any".into(), + after_ms: self.limits.io_timeout.as_millis(), + }); + break; + } + match reply_rx.recv_timeout(remaining) { + Ok(reply) => replies.push(reply), + Err(mpsc::RecvTimeoutError::Timeout) => { + failure = Some(MultiRailError::Timeout { + rail: "any".into(), + endpoint: "any".into(), + after_ms: self.limits.io_timeout.as_millis(), + }); + break; + } + Err(mpsc::RecvTimeoutError::Disconnected) => { + failure = Some(MultiRailError::WorkerPanic { + rail: "any".into(), + endpoint: "any".into(), + }); + break; + } + } + } + + // ---- verify replies, release in-flight accounting ---- + for reply in replies { + let rail = &self.rails[reply.rail_index]; + rail.stats.inflight_requests.fetch_sub(1, Ordering::Relaxed); + rail.stats + .inflight_bytes + .fetch_sub(reply.expected_bytes, Ordering::Relaxed); + match &reply.outcome { + Ok(Some(outcome)) => { + if outcome.bytes as u64 != reply.expected_bytes { + rail.stats.requests_err.fetch_add(1, Ordering::Relaxed); + failure = Some(MultiRailError::ByteCountMismatch { + rail: rail.config.device.clone(), + endpoint: reply.endpoint.to_string(), + expected: reply.expected_bytes, + actual: outcome.bytes as u64, + }); + failed_rail = Some(rail.index); + } else if outcome.num_chunks != reply.expected_chunks { + rail.stats.requests_err.fetch_add(1, Ordering::Relaxed); + failure = Some(MultiRailError::ChunkCountMismatch { + rail: rail.config.device.clone(), + endpoint: reply.endpoint.to_string(), + expected: reply.expected_chunks, + actual: outcome.num_chunks, + }); + failed_rail = Some(rail.index); + } else { + rail.stats.requests_ok.fetch_add(1, Ordering::Relaxed); + rail.stats + .bytes_read + .fetch_add(outcome.bytes as u64, Ordering::Relaxed); + verified_bytes += outcome.bytes as u64; + } + } + Ok(None) => { + // The server answers found=false both for absent + // objects and for descriptor mismatches; with a valid + // placement in hand this means a stale descriptor. + rail.stats.requests_err.fetch_add(1, Ordering::Relaxed); + failure = Some(MultiRailError::StaleDescriptor { + endpoint: reply.endpoint.to_string(), + }); + failed_rail = Some(rail.index); + } + Err(message) => { + rail.stats.requests_err.fetch_add(1, Ordering::Relaxed); + if reply.is_timeout() { + rail.stats.timeouts.fetch_add(1, Ordering::Relaxed); + failure = Some(MultiRailError::Timeout { + rail: rail.config.device.clone(), + endpoint: reply.endpoint.to_string(), + after_ms: self.limits.io_timeout.as_millis(), + }); + } else { + failure = Some(MultiRailError::TaskFailed { + rail: rail.config.device.clone(), + endpoint: reply.endpoint.to_string(), + source: message.clone(), + }); + } + failed_rail = Some(rail.index); + } + } + } + + if failure.is_some() { + break 'waves; + } + } + + // ---- failure path: quiesce participating connections ---- + if let Some(error) = failure { + if let Some(index) = failed_rail { + self.rails[index].mark_unhealthy(self.limits.rail_cooldown); + } + self.quiesce(participated.into_iter().collect()); + return Err(error); + } + + // Full-coverage invariant: the per-task byte checks must add up to + // the whole object, otherwise stripes went missing. + if verified_bytes != placement.expected_bytes { + return Err(MultiRailError::ByteCountMismatch { + rail: "aggregate".into(), + endpoint: "aggregate".into(), + expected: placement.expected_bytes, + actual: verified_bytes, + }); + } + + // ---- per-stripe checksum verification (when the server set them) ---- + for (stripe, info) in &placement.stripes { + if let Some(expected) = &info.checksum { + let start = info.offset as usize; + let end = start + info.length as usize; + if end > len { + return Err(MultiRailError::OutOfBoundsStripe { + stripe: *stripe, + stripe_count: placement.stripe_count, + }); + } + // SAFETY: `base..base+len` is the caller buffer, valid for + // this whole call; the slice below stays inside it. + let view: &[u8] = + unsafe { std::slice::from_raw_parts((base + start) as *const u8, info.length as usize) }; + let actual = format!("{:016x}", twox_hash::xxh3::hash64(view)); + if !expected.eq_ignore_ascii_case(&actual) { + return Err(MultiRailError::ChecksumMismatch { + stripe: *stripe, + expected: expected.clone(), + actual, + }); + } + } + } + + // ---- success path: synchronous MR eviction for non-sticky buffers ---- + if !sticky { + for key in &participated { + let ack_rx = { + let conns = self.conns.lock().unwrap(); + match conns.get(key) { + Some(entry) => { + let (ack_tx, ack_rx) = mpsc::channel(); + if entry + .tx + .send(Command::EvictRegistration { base, ack: ack_tx }) + .is_ok() + { + Some(ack_rx) + } else { + None + } + } + None => None, + } + }; + if let Some(ack_rx) = ack_rx { + let _ = ack_rx.recv_timeout(self.limits.io_timeout); + } + } + } + + Ok(placement.expected_bytes as usize) + } + + /// Wait until `bytes` more in-flight traffic fits the per-rail and total + /// budgets, or the deadline passes. + fn await_headroom(&self, rail: &Rail, bytes: u64, deadline: Instant) -> bool { + loop { + let rail_inflight = rail.stats.inflight_bytes.load(Ordering::Relaxed); + if rail_inflight + bytes <= self.limits.max_inflight_bytes_per_rail { + let total: u64 = self + .rails + .iter() + .map(|r| r.stats.inflight_bytes.load(Ordering::Relaxed)) + .sum(); + if total + bytes <= self.limits.max_inflight_bytes_total { + return true; + } + } + if Instant::now() >= deadline { + return false; + } + std::thread::sleep(Duration::from_millis(1)); + } + } + + /// Get an existing connection's command channel or spawn a worker for + /// (rail, endpoint), waiting for setup completion (bounded by `deadline`). + fn get_or_create_conn( + &self, + rail: &Arc, + endpoint: Arc, + deadline: Instant, + ) -> Result, MultiRailError> { + if let Some(entry) = self.conns.lock().unwrap().get(&(rail.index, Arc::clone(&endpoint))) { + return Ok(entry.tx.clone()); + } + + let (setup_tx, setup_rx) = mpsc::channel(); + let (command_tx, command_rx) = mpsc::channel(); + let worker_rail = Arc::clone(rail); + let worker_endpoint = Arc::clone(&endpoint); + let limits = self.limits.clone(); + let handle = std::thread::Builder::new() + .name(format!("mrail-{}-{}", rail.config.device, endpoint)) + .spawn(move || { + conn_worker(worker_rail, worker_endpoint, limits, command_rx, setup_tx) + }) + .map_err(|error| MultiRailError::TaskFailed { + rail: rail.config.device.clone(), + endpoint: endpoint.to_string(), + source: format!("spawn worker: {error}"), + })?; + + let remaining = deadline.saturating_duration_since(Instant::now()); + match setup_rx.recv_timeout(remaining) { + Ok(Ok(())) => {} + Ok(Err(reason)) => { + rail.mark_unhealthy(self.limits.rail_cooldown); + let _ = handle.join(); + return Err(MultiRailError::TaskFailed { + rail: rail.config.device.clone(), + endpoint: endpoint.to_string(), + source: format!("connect failed: {reason}"), + }); + } + Err(_) => { + rail.mark_unhealthy(self.limits.rail_cooldown); + let _ = command_tx.send(Command::Stop); + let _ = handle.join(); + return Err(MultiRailError::Timeout { + rail: rail.config.device.clone(), + endpoint: endpoint.to_string(), + after_ms: self.limits.io_timeout.as_millis(), + }); + } + } + + let mut conns = self.conns.lock().unwrap(); + // Another read may have created the same pair concurrently. + if let Some(existing) = conns.get(&(rail.index, Arc::clone(&endpoint))) { + let _ = command_tx.send(Command::Stop); + let _ = handle.join(); + return Ok(existing.tx.clone()); + } + conns.insert( + (rail.index, Arc::clone(&endpoint)), + ConnEntry { + rail_index: rail.index, + tx: command_tx.clone(), + handle, + }, + ); + Ok(command_tx) + } + + /// Stop and join the given connections, removing them from the map. Join + /// completion implies: worker loop exited → `RdmaClient` dropped → BYE + /// sent, QP destroyed, MRs deregistered. This is the quiesce barrier that + /// makes returning a failed read safe. + fn quiesce(&self, keys: Vec<(usize, Arc)>) { + let entries: Vec = { + let mut conns = self.conns.lock().unwrap(); + keys.into_iter().filter_map(|key| conns.remove(&key)).collect() + }; + for entry in &entries { + let _ = entry.tx.send(Command::Stop); + } + for entry in entries { + if entry.handle.join().is_ok() { + if let Some(rail) = self.rails.get(entry.rail_index) { + rail.stats + .connections_quiesced + .fetch_add(1, Ordering::Relaxed); + } + } + } + } +} + +impl Drop for MultiRailClient { + fn drop(&mut self) { + let keys: Vec<(usize, Arc)> = { + let conns = self.conns.lock().unwrap(); + conns.keys().cloned().collect() + }; + self.quiesce(keys); + } +} + +/// Verify the xxh3-64 checksum of one received stripe against the placement +/// checksum (lowercase hex, matching the server's `twox-hash` encoding). +pub fn verify_stripe_checksum( + buffer: &[u8], + stripe: u32, + chunk_size: u64, + expected_hex: &str, +) -> Result<(), MultiRailError> { + let offset = stripe as u64 * chunk_size; + let end = (offset + chunk_size).min(buffer.len() as u64); + if offset >= buffer.len() as u64 || end <= offset { + return Err(MultiRailError::OutOfBoundsStripe { + stripe, + stripe_count: stripe.saturating_add(1), + }); + } + let actual = twox_hash::xxh3::hash64(&buffer[offset as usize..end as usize]); + let actual_hex = format!("{actual:016x}"); + if !expected_hex.eq_ignore_ascii_case(&actual_hex) { + return Err(MultiRailError::ChecksumMismatch { + stripe, + expected: expected_hex.to_string(), + actual: actual_hex, + }); + } + Ok(()) +} + +// --------------------------------------------------------------------------- +// Tests (hardware-independent) +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + + fn descriptor(size: u64, stripe_count: u32, chunk_size: u64) -> pb::ObjectDescriptor { + pb::ObjectDescriptor { + key: Some(pb::ObjectKey { + namespace: "ns".into(), + object_key: "k".into(), + }), + object_handle: "h".into(), + object_generation: 1, + content_etag: "e".into(), + layout_version: 1, + size, + is_striped: stripe_count > 1, + stripe_count, + chunk_size, + } + } + + fn chunk(stripe: u32, endpoint: &str, length: u64, checksum: &str) -> pb::PlacementChunk { + pb::PlacementChunk { + stripe_index: stripe, + node_id: "node".into(), + grpc_endpoint: endpoint.into(), + rdma_endpoint: format!("{endpoint}:50053"), + device_id: 0, + storage_handle: "s".into(), + offset: stripe as u64 * 8, + length, + checksum: checksum.into(), + } + } + + fn test_rails(count: usize) -> Vec> { + (0..count) + .map(|index| { + Arc::new(Rail { + index, + config: RailConfig::new(format!("dev{index}")), + topology: RailTopology::default(), + stats: RailStats::default(), + healthy: AtomicBool::new(true), + unhealthy_until: Mutex::new(None), + gid_v4: None, + }) + }) + .collect() + } + + #[test] + fn rail_config_parses_optional_fields() { + assert_eq!(RailConfig::parse("mlx5_0").unwrap().port, 1); + let full = RailConfig::parse("irdma0:1:5:2").unwrap(); + assert_eq!(full.device, "irdma0"); + assert_eq!(full.port, 1); + assert_eq!(full.gid_index, 5); + assert_eq!(full.weight, 2); + assert!(RailConfig::parse(" ").is_none()); + } + + #[test] + fn placement_rejects_missing_and_duplicate_stripes() { + let desc = descriptor(24, 3, 8); + let chunks = vec![chunk(0, "10.0.0.1", 8, ""), chunk(2, "10.0.0.1", 8, "")]; + match validate_placement(&desc, &chunks) { + Err(MultiRailError::MissingStripes(stripes)) => assert_eq!(stripes, vec![1]), + other => panic!("expected missing stripes, got {other:?}"), + } + let dup = vec![ + chunk(0, "10.0.0.1", 8, ""), + chunk(0, "10.0.0.2", 8, ""), + chunk(1, "10.0.0.1", 8, ""), + chunk(2, "10.0.0.1", 8, ""), + ]; + match validate_placement(&desc, &dup) { + Err(MultiRailError::DuplicateStripes(stripes)) => assert_eq!(stripes, vec![0]), + other => panic!("expected duplicate stripes, got {other:?}"), + } + } + + #[test] + fn placement_rejects_out_of_bounds_and_bad_lengths() { + let desc = descriptor(24, 3, 8); + let oob = vec![ + chunk(0, "10.0.0.1", 8, ""), + chunk(1, "10.0.0.1", 8, ""), + chunk(2, "10.0.0.1", 8, ""), + chunk(3, "10.0.0.1", 8, ""), + ]; + assert!(matches!( + validate_placement(&desc, &oob), + Err(MultiRailError::OutOfBoundsStripe { stripe: 3, .. }) + )); + let bad_len = vec![ + chunk(0, "10.0.0.1", 8, ""), + chunk(1, "10.0.0.1", 7, ""), + chunk(2, "10.0.0.1", 8, ""), + ]; + assert!(matches!( + validate_placement(&desc, &bad_len), + Err(MultiRailError::InvalidPlacement(_)) + )); + // Last stripe is short: a 20-byte object with 8-byte chunks ends in 4. + let desc2 = descriptor(20, 3, 8); + let ok = vec![ + chunk(0, "10.0.0.1", 8, ""), + chunk(1, "10.0.0.1", 8, ""), + chunk(2, "10.0.0.1", 4, ""), + ]; + let validated = validate_placement(&desc2, &ok).expect("short last stripe accepted"); + assert_eq!(validated.expected_bytes, 20); + } + + #[test] + fn placement_accepts_non_striped_single_chunk() { + let desc = descriptor(100, 0, 0); + let chunks = vec![chunk(0, "10.0.0.1", 0, "")]; + let validated = validate_placement(&desc, &chunks).expect("valid"); + assert_eq!(validated.expected_bytes, 100); + assert_eq!(validated.stripe_count, 1); + } + + #[test] + fn plan_balances_stripes_over_rails_by_bytes() { + let desc = descriptor(64, 8, 8); + let chunks: Vec<_> = (0..8).map(|i| chunk(i, "10.0.0.1", 8, "")).collect(); + let placement = validate_placement(&desc, &chunks).unwrap(); + let rails = test_rails(2); + let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded).unwrap(); + assert_eq!(plan.task_count(), 2); + let total: u64 = plan.tasks.iter().map(|t| t.bytes).sum(); + assert_eq!(total, 64); + // 8 equal stripes over 2 rails → 4/4 per rail. + for task in &plan.tasks { + assert_eq!(task.bytes, 32); + assert_eq!(task.stripes.len(), 4); + } + // Disjoint stripe sets covering 0..8. + let mut seen: HashSet = HashSet::new(); + for task in &plan.tasks { + for stripe in &task.stripes { + assert!(seen.insert(*stripe), "stripe {stripe} assigned twice"); + } + } + assert_eq!(seen.len(), 8); + } + + #[test] + fn plan_respects_relative_weights() { + let desc = descriptor(80, 10, 8); + let chunks: Vec<_> = (0..10).map(|i| chunk(i, "10.0.0.1", 8, "")).collect(); + let placement = validate_placement(&desc, &chunks).unwrap(); + let mut rails = test_rails(2); + rails[1] = Arc::new(Rail { + index: 1, + config: RailConfig::new("dev1").with_weight(4), + topology: RailTopology::default(), + stats: RailStats::default(), + healthy: AtomicBool::new(true), + unhealthy_until: Mutex::new(None), + gid_v4: None, + }); + let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded).unwrap(); + let rail0: u64 = plan + .tasks + .iter() + .filter(|t| t.rail_index == 0) + .map(|t| t.bytes) + .sum(); + let rail1: u64 = plan + .tasks + .iter() + .filter(|t| t.rail_index == 1) + .map(|t| t.bytes) + .sum(); + // 1:4 weights over 10 stripes of 8 bytes: 16 vs 64. + assert_eq!((rail0, rail1), (16, 64)); + } + + #[test] + fn rail_endpoint_whitelist_restricts_and_errors() { + let desc = descriptor(16, 2, 8); + let chunks = vec![chunk(0, "10.0.0.1", 8, ""), chunk(1, "10.0.0.2", 8, "")]; + let placement = validate_placement(&desc, &chunks).unwrap(); + let mut rails = test_rails(2); + rails[0] = Arc::new(Rail { + index: 0, + config: RailConfig::new("dev0").with_endpoints(vec!["10.0.0.2".into()]), + topology: RailTopology::default(), + stats: RailStats::default(), + healthy: AtomicBool::new(true), + unhealthy_until: Mutex::new(None), + gid_v4: None, + }); + let plan = + build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded).expect("plan builds"); + // dev0 pinned to 10.0.0.2; 10.0.0.1 falls to dev1. + assert_eq!(plan.tasks.len(), 2); + for task in &plan.tasks { + let expected_rail = if task.endpoint.contains("10.0.0.2") { 0 } else { 1 }; + assert_eq!(task.rail_index, expected_rail); + } + // No rail allowed for 10.0.0.1 → typed planning error. + let mut strict = test_rails(1); + strict[0] = Arc::new(Rail { + index: 0, + config: RailConfig::new("dev0").with_endpoints(vec!["10.0.0.2".into()]), + topology: RailTopology::default(), + stats: RailStats::default(), + healthy: AtomicBool::new(true), + unhealthy_until: Mutex::new(None), + gid_v4: None, + }); + assert!(build_plan(&placement, &strict, RailSelectPolicy::LeastLoaded).is_err()); + } + + #[test] + fn plan_pins_endpoint_with_affinity_policy() { + let desc = descriptor(64, 8, 8); + let chunks: Vec<_> = (0..8).map(|i| chunk(i, "10.0.0.1", 8, "")).collect(); + let placement = validate_placement(&desc, &chunks).unwrap(); + let rails = test_rails(2); + let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity).unwrap(); + assert_eq!(plan.task_count(), 1); + assert_eq!(plan.tasks[0].bytes, 64); + assert_eq!(plan.tasks[0].stripes.len(), 8); + } + + #[test] + fn affinity_prefers_same_subnet_rail() { + let desc = descriptor(16, 2, 8); + let chunks = vec![chunk(0, "10.0.0.9", 8, ""), chunk(1, "10.0.0.9", 8, "")]; + let placement = validate_placement(&desc, &chunks).unwrap(); + let mut rails = test_rails(2); + rails[1] = Arc::new(Rail { + index: 1, + config: RailConfig::new("dev1"), + topology: RailTopology::default(), + stats: RailStats::default(), + healthy: AtomicBool::new(true), + unhealthy_until: Mutex::new(None), + gid_v4: Some("10.0.0.5".parse().unwrap()), + }); + let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity).unwrap(); + assert_eq!(plan.tasks.len(), 1); + assert_eq!(plan.tasks[0].rail_index, 1); + } + + #[test] + fn waves_respect_connection_caps() { + let tasks: Vec = (0..6) + .map(|i| TaskSpec { + rail_index: i % 2, + endpoint: Arc::from(format!("10.0.0.{i}:50053")), + stripes: vec![i as u32], + bytes: 8, + }) + .collect(); + let limits = RailLimits { + max_connections_per_rail: 2, + max_connections_total: 4, + ..RailLimits::default() + }; + let waves = plan_waves(tasks, &limits); + assert!(waves.len() >= 2); + for wave in &waves { + assert!(wave.len() <= limits.max_connections_total); + let mut per_rail: HashMap = HashMap::new(); + for task in wave { + *per_rail.entry(task.rail_index).or_insert(0) += 1; + } + for count in per_rail.values() { + assert!(*count <= limits.max_connections_per_rail); + } + } + } + + #[test] + fn gid_ipv4_mapping() { + let mut raw = [0u8; 16]; + raw[10] = 0xff; + raw[11] = 0xff; + raw[12..16].copy_from_slice(&[10, 0, 0, 7]); + assert_eq!( + gid_to_v4(raw).map(|ip| ip.to_string()), + Some("10.0.0.7".into()) + ); + assert!(gid_to_v4([0u8; 16]).is_none()); + } + + #[test] + fn endpoint_host_extraction() { + assert_eq!( + endpoint_v4("10.1.2.3:50053").map(|ip| ip.to_string()), + Some("10.1.2.3".into()) + ); + assert_eq!(endpoint_v4("not-a-host:50053"), None); + } + + #[test] + fn checksum_verification_matches_server_encoding() { + let mut buffer = vec![0u8; 16]; + buffer[..8].copy_from_slice(b"stripe0!"); + buffer[8..].copy_from_slice(b"stripe1!"); + let expected0 = format!("{:016x}", twox_hash::xxh3::hash64(&buffer[..8])); + assert!(verify_stripe_checksum(&buffer, 0, 8, &expected0).is_ok()); + let wrong = format!("{:016x}", twox_hash::xxh3::hash64(&buffer[8..])); + assert!(matches!( + verify_stripe_checksum(&buffer, 0, 8, &wrong), + Err(MultiRailError::ChecksumMismatch { .. }) + )); + // Out-of-bounds stripe. + assert!(matches!( + verify_stripe_checksum(&buffer, 2, 8, &expected0), + Err(MultiRailError::OutOfBoundsStripe { .. }) + )); + } + + #[test] + fn topology_parses_sysfs_shape() { + let dir = tempfile::tempdir().unwrap(); + let dev = dir.path().join("rxe0/device"); + std::fs::create_dir_all(&dev).unwrap(); + std::fs::write(dev.join("numa_node"), "1\n").unwrap(); + let topology = topology_from(dir.path(), "rxe0"); + assert_eq!(topology.numa_node, 1); + assert!(!topology.pci_slot.is_empty()); + } + + #[test] + fn cooldown_expires_and_rail_recovers() { + let rails = test_rails(1); + rails[0].mark_unhealthy(Duration::from_millis(20)); + assert!(!rails[0].selectable()); + std::thread::sleep(Duration::from_millis(30)); + assert!(rails[0].selectable()); + } + + #[test] + fn policy_parsing_round_trip() { + assert_eq!( + RailSelectPolicy::parse("endpoint-affinity"), + Some(RailSelectPolicy::EndpointAffinity) + ); + assert_eq!( + RailSelectPolicy::parse("least-loaded"), + Some(RailSelectPolicy::LeastLoaded) + ); + assert_eq!(RailSelectPolicy::parse("bogus"), None); + } + + #[test] + fn expected_task_outcome_matches_layout() { + let desc = descriptor(20, 3, 8); + let (bytes, chunks) = expected_task_outcome(&desc, &[0, 2]); + assert_eq!((bytes, chunks), (12, 2)); + let (bytes, chunks) = expected_task_outcome(&desc, &[]); + assert_eq!((bytes, chunks), (20, 1)); + } +} diff --git a/kv-service/client-rs/src/rdma.rs b/kv-service/client-rs/src/rdma.rs index a7532d0..964e5ae 100644 --- a/kv-service/client-rs/src/rdma.rs +++ b/kv-service/client-rs/src/rdma.rs @@ -16,7 +16,7 @@ use rdma_sys::*; use std::ffi::{c_void, CStr}; use std::io::{Read, Write}; use std::marker::PhantomData; -use std::net::TcpStream; +use std::net::{TcpStream, ToSocketAddrs}; use std::ptr::{self, NonNull}; use std::sync::Arc; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; @@ -59,6 +59,12 @@ pub struct RdmaClientConfig { pub port: u8, /// GID index used to construct the RoCE address handle. pub gid_index: u8, + /// Per-operation timeout applied to the TCP control channel. A blocked or + /// dead peer surfaces as an io error instead of an indefinite hang, which + /// is what lets multi-rail callers fail a rail safely. + pub io_timeout: Option, + /// Bound on the TCP connect phase itself. + pub connect_timeout: Option, } impl RdmaClientConfig { @@ -69,6 +75,8 @@ impl RdmaClientConfig { device: device.into(), port: 1, gid_index: 3, + io_timeout: None, + connect_timeout: None, } } @@ -83,6 +91,92 @@ impl RdmaClientConfig { self.gid_index = gid_index; self } + + /// Set the per-operation TCP control-channel timeout. + pub fn with_io_timeout(mut self, timeout: Duration) -> Self { + self.io_timeout = Some(timeout); + self + } + + /// Set the TCP connect timeout. + pub fn with_connect_timeout(mut self, timeout: Duration) -> Self { + self.connect_timeout = Some(timeout); + self + } +} + +/// Outcome of a descriptor GET: how many bytes the server placed in the +/// destination and how many stripes/chunks it served. `num_chunks` comes from +/// the tag-3 response body and lets callers detect missing or duplicated +/// stripes server-side. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct GetOutcome { + pub bytes: usize, + pub num_chunks: u32, +} + +/// Enumerate the verbs devices present on this host, by name. +pub fn list_device_names() -> Vec { + unsafe { + let mut count = 0i32; + let devices = ibv_get_device_list(&mut count); + if devices.is_null() { + return Vec::new(); + } + let mut names = Vec::with_capacity(count.max(0) as usize); + for index in 0..count { + let device = *devices.offset(index as isize); + if device.is_null() { + continue; + } + names.push( + CStr::from_ptr(ibv_get_device_name(device)) + .to_string_lossy() + .into_owned(), + ); + } + ibv_free_device_list(devices); + names + } +} + +/// Read one GID of a device without keeping a context open. Used by callers +/// that need the local address (e.g. for subnet-affinity rail selection). +pub fn query_gid_raw(device: &str, port: u8, gid_index: u8) -> Option<[u8; 16]> { + unsafe { + let mut count = 0i32; + let devices = ibv_get_device_list(&mut count); + if devices.is_null() { + return None; + } + let mut selected = ptr::null_mut(); + for index in 0..count { + let candidate = *devices.offset(index as isize); + if candidate.is_null() { + continue; + } + if CStr::from_ptr(ibv_get_device_name(candidate)).to_string_lossy() == device { + selected = candidate; + break; + } + } + if selected.is_null() { + ibv_free_device_list(devices); + return None; + } + let context = ibv_open_device(selected); + ibv_free_device_list(devices); + if context.is_null() { + return None; + } + let mut gid: ibv_gid = std::mem::zeroed(); + let rc = ibv_query_gid(context, port, gid_index as i32, &mut gid); + ibv_close_device(context); + if rc != 0 { + return None; + } + Some(gid.raw) + } } /// Result of an RDMA read. `None` means the object did not exist. @@ -296,13 +390,20 @@ impl RdmaClient { } let result = (|| -> Result { - let mut stream = TcpStream::connect(&config.endpoint) - .with_context(|| format!("connect RDMA control endpoint {}", config.endpoint))?; + let mut stream = tcp_connect(&config)?; // 控制面小包必须即时发出: Nagle + delayed-ACK 在多连接并发时会给每个 // 请求注入 ~40-200ms 延迟 (数据面 RDMA WRITE 不经 TCP, 不受影响). stream .set_nodelay(true) .context("set TCP_NODELAY on RDMA control stream")?; + if let Some(timeout) = config.io_timeout { + stream + .set_read_timeout(Some(timeout)) + .context("set RDMA control read timeout")?; + stream + .set_write_timeout(Some(timeout)) + .context("set RDMA control write timeout")?; + } let local = QpInfo { qpn: unsafe { (*qp.as_ptr()).qp_num }, psn: random_psn(), @@ -509,6 +610,20 @@ impl RdmaClient { view: BufferView, offset: usize, ) -> Result { + self.get_descriptor_stripes_into_view_detailed(descriptor, stripes, view, offset) + .map(|outcome| outcome.map(|outcome| outcome.bytes)) + } + + /// Detailed variant of [`Self::get_descriptor_stripes_into_view`] that also + /// reports how many stripes the server actually served, enabling + /// missing/duplicate-stripe detection on the caller side. + pub fn get_descriptor_stripes_into_view_detailed( + &mut self, + descriptor: &pb::ObjectDescriptor, + stripes: &[u32], + view: BufferView, + offset: usize, + ) -> Result> { let key = descriptor .key .as_ref() @@ -532,7 +647,17 @@ impl RdmaClient { } self.stream.write_all(&request)?; self.stream.flush()?; - read_get_response(&mut self.stream) + read_get_response_detailed(&mut self.stream) + } + + /// Drop every cached registration whose base pointer is `base`, + /// deregistering the memory regions. Returns how many entries were + /// evicted. Multi-rail callers use this to quiesce registrations of a + /// caller buffer synchronously before returning it. + pub fn evict_registrations_for(&mut self, base: usize) -> usize { + let before = self.mr_cache.len(); + self.mr_cache.retain(|((ptr, _), _)| *ptr != base); + before - self.mr_cache.len() } /// Stripe-subset GET with a scatter destination list (wire tag 15): the @@ -916,6 +1041,22 @@ fn random_psn() -> u32 { & 0x00ff_ffff } +fn tcp_connect(config: &RdmaClientConfig) -> Result { + let mut addresses = config + .endpoint + .to_socket_addrs() + .with_context(|| format!("resolve RDMA control endpoint {}", config.endpoint))?; + let address = addresses + .next() + .ok_or_else(|| anyhow!("RDMA endpoint resolved to no address: {}", config.endpoint))?; + match config.connect_timeout { + Some(timeout) => TcpStream::connect_timeout(&address, timeout) + .with_context(|| format!("connect RDMA control endpoint {address}")), + None => TcpStream::connect(address) + .with_context(|| format!("connect RDMA control endpoint {address}")), + } +} + fn create_qp(resources: &RdmaResources) -> Result> { let mut attr = ibv_qp_init_attr { qp_context: ptr::null_mut(), @@ -1137,6 +1278,10 @@ fn read_string(stream: &mut TcpStream, field: &str) -> Result { } fn read_get_response(stream: &mut TcpStream) -> Result { + Ok(read_get_response_detailed(stream)?.map(|outcome| outcome.bytes)) +} + +fn read_get_response_detailed(stream: &mut TcpStream) -> Result> { let mut tag = [0u8; 1]; stream.read_exact(&mut tag)?; if tag[0] != MSG_GET_RESP { @@ -1149,7 +1294,8 @@ fn read_get_response(stream: &mut TcpStream) -> Result { } let bytes = u64::from_le_bytes(body[1..9].try_into().expect("fixed get response length")); let bytes = usize::try_from(bytes).map_err(|_| anyhow!("RDMA read size exceeds usize"))?; - Ok(Some(bytes)) + let num_chunks = u32::from_le_bytes(body[9..13].try_into().expect("fixed chunk count")); + Ok(Some(GetOutcome { bytes, num_chunks })) } struct PutReady { @@ -1288,7 +1434,7 @@ mod tests { fn request_rejects_oversized_wire_string() { let key = "x".repeat(u16::MAX as usize + 1); assert!(build_get_request(&key, 1, 2, 3).is_err()); - assert!(build_put_request(MSG_PUT_REQ, &key, 3).is_err()); + assert!(build_put_request(MSG_PUT_REQ, &key, 3, 0).is_err()); } #[test] diff --git a/kv-service/client-rs/tests/multirail_e2e.rs b/kv-service/client-rs/tests/multirail_e2e.rs new file mode 100644 index 0000000..b2408e4 --- /dev/null +++ b/kv-service/client-rs/tests/multirail_e2e.rs @@ -0,0 +1,301 @@ +//! Hardware-gated end-to-end coverage for the multi-rail RDMA read path. +//! +//! Run explicitly on a host with an RDMA-enabled ContextStore server that +//! exposes at least two RDMA listeners (e.g. Soft-RoCE): +//! ```text +//! CS_MR_COORDINATOR=http://127.0.0.1:50051 \ +//! CS_MR_RAILS=rxe0,rxe1 \ +//! cargo test --manifest-path kv-service/client-rs/Cargo.toml --features rdma \ +//! --test multirail_e2e -- --ignored --nocapture +//! ``` + +#![cfg(feature = "rdma")] + +use contextstore_client_rs::multirail::{ + MultiRailClient, MultiRailError, RailConfig, RailLimits, +}; +use contextstore_client_rs::{KvClient, ObjectLookup}; +use prost::bytes::Bytes; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +fn env_or(name: &str, default: &str) -> String { + std::env::var(name).unwrap_or_else(|_| default.to_string()) +} + +fn rails_from_env() -> Vec { + let pins = pin_map_from_env(); + env_or("CS_MR_RAILS", "rxe0,rxe1") + .split(',') + .filter_map(RailConfig::parse) + .map(|mut rail| { + if let Some(endpoints) = pins.get(&rail.device) { + rail.endpoints = endpoints.clone(); + } + rail + }) + .collect() +} + +/// Rails with endpoint whitelists cleared (single-rail reference reads and +/// failure injection must be able to reach any endpoint). +fn unpinned(rails: &[RailConfig]) -> Vec { + rails + .iter() + .cloned() + .map(|mut rail| { + rail.endpoints.clear(); + rail + }) + .collect() +} + +/// `CS_MR_PIN="rxe0=host[:port][;host...],rxe1=..."` → device → endpoints. +fn pin_map_from_env() -> std::collections::HashMap> { + let mut map = std::collections::HashMap::new(); + for spec in env_or("CS_MR_PIN", "").split(',') { + if let Some((device, targets)) = spec.split_once('=') { + let endpoints: Vec = targets + .split(';') + .map(|t| t.trim().to_string()) + .filter(|t| !t.is_empty()) + .collect(); + if !endpoints.is_empty() { + map.insert(device.trim().to_string(), endpoints); + } + } + } + map +} + +/// Rewrite the placement's per-stripe RDMA endpoints round-robin across +/// `CS_MR_ALTERNATE_ENDPOINTS` (comma-separated `host:port` list). Used on +/// single-host testbeds where one node owns several NIC listeners. +fn remap_lookup_endpoints(lookup: &ObjectLookup) -> ObjectLookup { + let list = env_or("CS_MR_ALTERNATE_ENDPOINTS", ""); + if list.is_empty() { + return lookup.clone(); + } + let endpoints: Vec = list.split(',').map(|s| s.trim().to_string()).collect(); + let mut remapped = lookup.clone(); + if let Some(placement) = remapped.placement.as_mut() { + for (index, chunk) in placement.chunks.iter_mut().enumerate() { + chunk.rdma_endpoint = endpoints[index % endpoints.len()].clone(); + } + } + remapped +} + +fn unique_key(prefix: &str) -> String { + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .expect("system clock before Unix epoch") + .as_nanos(); + format!("{prefix}-{nanos}") +} + +fn pattern_word(word_index: u64) -> u64 { + word_index + .wrapping_mul(0x9E37_79B9_7F4A_7C15) + .rotate_left(23) + ^ word_index +} + +fn fill_pattern(buffer: &mut [u8]) { + for (index, chunk) in buffer.chunks_exact_mut(8).enumerate() { + chunk.copy_from_slice(&pattern_word(index as u64).to_le_bytes()); + } +} + +fn verify_pattern(buffer: &[u8]) -> bool { + buffer + .chunks_exact(8) + .enumerate() + .all(|(index, chunk)| chunk == pattern_word(index as u64).to_le_bytes()) +} + +struct Fixture { + runtime: tokio::runtime::Runtime, + namespace: String, + key: String, + lookup: ObjectLookup, + size: usize, +} + +/// Seed a striped object over gRPC and look it up. The ContextStore server +/// stripes large objects across nodes/NICs per its placement policy; a +/// single-node two-NIC server still produces a striped layout with both +/// endpoints when the object exceeds the stripe threshold. +fn seed(namespace: &str, key: &str, size_mb: usize) -> Fixture { + let runtime = tokio::runtime::Runtime::new().expect("create tokio runtime"); + let size = size_mb * 1024 * 1024; + let mut payload = vec![0u8; size]; + fill_pattern(&mut payload); + let coordinator = env_or("CS_MR_COORDINATOR", "http://127.0.0.1:50051"); + let lookup = runtime + .block_on(async { + let mut client = KvClient::connect(coordinator) + .await + .expect("connect gRPC coordinator"); + let big = Bytes::from(payload); + let chunk = 4 * 1024 * 1024; + let mut segments = Vec::new(); + for offset in (0..size).step_by(chunk) { + segments.push(big.slice(offset..(offset + chunk).min(size))); + } + client + .put_stream_chunks(namespace, key, segments) + .await + .expect("seed object over gRPC"); + client + .lookup_object(namespace, key) + .await + .expect("lookup seeded object") + .expect("seeded object present") + }); + Fixture { + runtime, + namespace: namespace.to_string(), + key: key.to_string(), + size: usize::try_from(lookup.descriptor.size).expect("size fits usize"), + lookup, + } +} + +fn limits() -> RailLimits { + RailLimits { + io_timeout: Duration::from_secs(20), + rail_cooldown: Duration::from_millis(500), + ..RailLimits::default() + } +} + +#[test] +#[ignore = "requires an RDMA-enabled ContextStore server with 2+ listeners"] +fn multirail_read_returns_identical_bytes_across_rail_counts() { + let fixture = seed("multirail-e2e", &unique_key("dual"), 128); + assert!( + fixture.lookup.placement.is_some(), + "lookup returned no placement" + ); + let rails = rails_from_env(); + assert!(rails.len() >= 2, "CS_MR_RAILS must list at least two rails"); + let lookup = remap_lookup_endpoints(&fixture.lookup); + + // Dual-rail read. + let dual = MultiRailClient::new(rails.clone()) + .expect("create dual-rail client") + .with_limits(limits()); + let mut buffer = vec![0xA5u8; fixture.size]; + let bytes = dual + .read_lookup_into(&lookup, &mut buffer) + .expect("dual-rail read"); + assert_eq!(bytes, fixture.size); + assert!(verify_pattern(&buffer), "dual-rail content mismatch"); + + // Multi-rail metrics must show traffic on more than one rail. + let snapshots = dual.rails_snapshot(); + let active = snapshots.iter().filter(|s| s.bytes_read > 0).count(); + assert!( + active >= 2, + "expected bytes on >=2 rails, got {active}: {snapshots:?}" + ); + + // Single-rail read of the same object must match byte-for-byte + // (compatibility: the safe-API buffer is reused across clients). On + // cross-wired testbeds a lone rail may only reach its own-side listener, + // so the reference reads the original (un-remapped) placement. + let single = MultiRailClient::new(unpinned(&rails[..1])) + .expect("create single-rail client") + .with_limits(limits()); + let mut single_buffer = vec![0x5Au8; fixture.size]; + let single_bytes = single + .read_lookup_into(&fixture.lookup, &mut single_buffer) + .expect("single-rail read"); + assert_eq!(single_bytes, fixture.size); + assert_eq!(single_buffer, buffer, "rail counts changed content"); +} + +#[test] +#[ignore = "requires an RDMA-enabled ContextStore server with 2+ listeners"] +fn rail_failure_fails_safely_and_next_read_recovers() { + let fixture = seed("multirail-e2e", &unique_key("fail"), 128); + let rails = rails_from_env(); + let dead_endpoint = env_or("CS_MR_DEAD_ENDPOINT", "127.0.0.1:59999"); + + // Sabotage the placement: point every stripe at a dead endpoint. The + // read must fail (safe failure — no transparent retry) and return + // without leaving in-flight work behind. Rails are unpinned here so the + // failure surfaces as a transport error, not a planning error. + let mut dead_chunks = fixture + .lookup + .placement + .as_ref() + .expect("placement") + .chunks + .clone(); + for chunk in &mut dead_chunks { + chunk.rdma_endpoint = dead_endpoint.clone(); + } + let client = MultiRailClient::new(unpinned(&rails)) + .expect("create client") + .with_limits(limits()); + let mut buffer = vec![0xA5u8; fixture.size]; + match client.read_object_into(&fixture.lookup.descriptor, &dead_chunks, &mut buffer) { + Err(error @ (MultiRailError::TaskFailed { .. } | MultiRailError::Timeout { .. })) => { + let text = error.to_string(); + assert!( + text.contains(dead_endpoint.as_str()), + "error should mention the dead endpoint: {text}" + ); + } + other => panic!("expected safe failure, got {other:?}"), + } + + // Path recovery: the very same buffer, still poisoned, is served fine by + // a client whose rails point at the live endpoints. A late RDMA WRITE + // from the failed attempt would corrupt this verification. + let recovered = MultiRailClient::new(rails) + .expect("create recovery client") + .with_limits(limits()); + let bytes = recovered + .read_lookup_into(&remap_lookup_endpoints(&fixture.lookup), &mut buffer) + .expect("recovery read after failure"); + assert_eq!(bytes, fixture.size); + assert!(verify_pattern(&buffer), "recovered content mismatch"); +} + +#[test] +#[ignore = "requires an RDMA-enabled ContextStore server with 2+ listeners"] +fn stale_descriptor_is_reported_as_relookup() { + let fixture = seed("multirail-e2e", &unique_key("stale"), 128); + let coordinator = env_or("CS_MR_COORDINATOR", "http://127.0.0.1:50051"); + + // Rewrite the object (new generation) through gRPC after the lookup. + let mut payload = vec![0u8; fixture.size]; + fill_pattern(&mut payload); + payload[0] ^= 0xFF; // different content → different etag/generation + fixture + .runtime + .block_on(async { + let mut client = KvClient::connect(coordinator) + .await + .expect("connect coordinator"); + client + .delete(&fixture.namespace, &fixture.key) + .await + .expect("delete old version"); + client + .put(&fixture.namespace, &fixture.key, payload) + .await + .expect("rewrite object"); + }); + + let client = MultiRailClient::new(rails_from_env()) + .expect("create client") + .with_limits(limits()); + let mut buffer = vec![0u8; fixture.size]; + match client.read_lookup_into(&remap_lookup_endpoints(&fixture.lookup), &mut buffer) { + Err(MultiRailError::StaleDescriptor { .. }) => {} + other => panic!("expected StaleDescriptor, got {other:?}"), + } +} diff --git a/kv-service/server/src/lib.rs b/kv-service/server/src/lib.rs index 6155974..7e39e52 100644 --- a/kv-service/server/src/lib.rs +++ b/kv-service/server/src/lib.rs @@ -9,6 +9,11 @@ //! - `metadata` : Prefix Index + Block Allocator //! - `config` : configuration loading +// tonic-build's generated service stubs return Result<_, tonic::Status>; +// newer clippy versions flag the Status size on that generated code. The +// generated file is vendored from the proto, so silence the lint crate-wide. +#![allow(clippy::result_large_err)] + pub mod api; pub mod config; pub mod error; From aaff611308dea8a3f4a1eb8c03cb5bff1634f679 Mon Sep 17 00:00:00 2001 From: wzu Date: Sat, 26 Sep 2026 10:17:24 +0800 Subject: [PATCH 2/7] multirail: path-MTU config, intra-rail connection splitting, testbed hardening - RC path MTU configurable end to end (RdmaClientConfig with_path_mtu, RailConfig 5th spec field, server CS_RDMA_PATH_MTU); default stays 1024. Both peers must agree - the HELLO exchange does not carry it. - RailLimits task_max_stripes splits per-rail stripes over several connections (intra-rail concurrency / queue depth); exercised by the e2e suite via CS_MR_TASK_MAX_STRIPES. - Bench: --sticky (cached registrations), --task-max-stripes. - Testbed moved to a conflict-free 192.168.250.0/24 veth pair with an up/down setup script; full netns isolation is blocked by an rxe in-netns transport bug (vanilla upstream PUT flushes with WR_FLUSH_ERR), documented in the notes. --- .../client-rs/src/bin/multirail_bench.rs | 35 +++++++++- kv-service/client-rs/src/multirail.rs | 70 ++++++++++++++----- kv-service/client-rs/src/rdma.rs | 25 ++++++- kv-service/client-rs/tests/multirail_e2e.rs | 12 ++++ kv-service/server/src/rdma/qp.rs | 10 ++- kv-service/server/src/rdma/server.rs | 19 ++++- 6 files changed, 146 insertions(+), 25 deletions(-) diff --git a/kv-service/client-rs/src/bin/multirail_bench.rs b/kv-service/client-rs/src/bin/multirail_bench.rs index 45d595c..4dcb29c 100644 --- a/kv-service/client-rs/src/bin/multirail_bench.rs +++ b/kv-service/client-rs/src/bin/multirail_bench.rs @@ -49,6 +49,18 @@ struct Args { /// Stripe→rail policy: least-loaded | endpoint-affinity | rr #[arg(long, default_value = "least-loaded")] policy: String, + /// RC path MTU for all rails (bytes). 4096 on jumbo-frame fabrics. + #[arg(long, default_value = "1024")] + qp_mtu: u16, + /// Keep per-rail registrations cached across iterations (pinned-buffer + /// fast path; skips ~ibv_reg_mr of the whole buffer every read). + #[arg(long, default_value_t = false)] + sticky: bool, + /// Max stripes per task: splits an endpoint's stripes over several + /// connections per rail (intra-rail concurrency). 0 = one task per + /// (rail, endpoint). + #[arg(long, default_value = "0")] + task_max_stripes: usize, #[arg(long, default_value = "5")] iters: usize, /// Destination buffer size in MiB (>= object size). @@ -190,6 +202,7 @@ fn run_client( .with_context(|| format!("unknown policy '{}'", args.policy))?; let limits = RailLimits { io_timeout: std::time::Duration::from_secs(args.io_timeout_secs), + task_max_stripes: args.task_max_stripes, ..RailLimits::default() }; let client = MultiRailClient::new(rails)?.with_limits(limits).with_policy(policy); @@ -214,9 +227,22 @@ fn run_client( // Poison the buffer so a missing stripe cannot slip through. buffer.as_mut().iter_mut().for_each(|b| *b = 0xA5); let started = Instant::now(); - let bytes = client - .read_lookup_into(lookup, buffer.as_mut()) - .with_context(|| format!("[{label}] iteration {iteration} failed"))?; + let bytes = if args.sticky { + // SAFETY: the AlignedBuffer outlives the client and is never + // freed or reused while reads run. + unsafe { + client.read_object_into_raw( + &lookup.descriptor, + &lookup.placement.as_ref().map(|p| p.chunks.clone()).unwrap_or_default(), + buffer.ptr, + buffer.len, + true, + ) + } + } else { + client.read_lookup_into(lookup, buffer.as_mut()) + } + .with_context(|| format!("[{label}] iteration {iteration} failed"))?; latencies.push(started.elapsed()); last_len = bytes; if bytes != object_size { @@ -298,6 +324,9 @@ fn main() -> Result<()> { if let Some(endpoints) = pins.get(&rail.device) { rail.endpoints = endpoints.clone(); } + if args.qp_mtu != 1024 { + rail.mtu = args.qp_mtu; + } rail }) }) diff --git a/kv-service/client-rs/src/multirail.rs b/kv-service/client-rs/src/multirail.rs index 8d37df3..0143ad4 100644 --- a/kv-service/client-rs/src/multirail.rs +++ b/kv-service/client-rs/src/multirail.rs @@ -89,6 +89,9 @@ pub struct RailConfig { pub gid_index: u8, /// Relative share of stripes this rail receives (load-balancing weight). pub weight: u32, + /// RC path MTU in bytes; 4096 on jumbo-frame fabrics, 1024 (default) on + /// standard 1500-byte networks. + pub mtu: u16, /// Optional endpoint whitelist (`host:port` or bare `host`). Empty means /// the rail may serve any endpoint. Rail-optimized fabrics pin rail k to /// fabric k; soft-RoCE cross-wired testbeds use it the same way. @@ -102,6 +105,7 @@ impl RailConfig { port: 1, gid_index: 3, weight: 1, + mtu: 1024, endpoints: Vec::new(), } } @@ -121,6 +125,12 @@ impl RailConfig { self } + /// Set the RC path MTU (bytes; 4096 for jumbo-frame fabrics). + pub fn with_mtu(mut self, bytes: u16) -> Self { + self.mtu = bytes; + self + } + /// Restrict this rail to the given endpoints (`host:port` or bare host). pub fn with_endpoints(mut self, endpoints: Vec) -> Self { self.endpoints = endpoints; @@ -139,7 +149,7 @@ impl RailConfig { .any(|allowed| allowed == endpoint || *allowed == host) } - /// Parse `device[:port[:gid[:weight]]]` (portions left out keep defaults). + /// Parse `device[:port[:gid[:weight[:mtu]]]]` (omitted parts keep defaults). pub fn parse(spec: &str) -> Option { let mut parts = spec.split(':'); let device = parts.next()?.trim(); @@ -156,6 +166,9 @@ impl RailConfig { if let Some(weight) = parts.next().and_then(|w| w.trim().parse().ok()) { config = config.with_weight(weight); } + if let Some(mtu) = parts.next().and_then(|m| m.trim().parse().ok()) { + config = config.with_mtu(mtu); + } Some(config) } } @@ -171,6 +184,10 @@ pub struct RailLimits { pub max_inflight_bytes_per_rail: u64, /// In-flight bytes across all rails before dispatch blocks. pub max_inflight_bytes_total: u64, + /// Cap on stripes per task: an endpoint's stripes assigned to one rail + /// are split into tasks of at most this many stripes, each on its own + /// connection (per-rail concurrency / queue depth). 0 = unlimited. + pub task_max_stripes: usize, /// Per-operation TCP control-channel timeout; also bounds connect. pub io_timeout: Duration, /// How long a failed rail is skipped before it is retried. @@ -184,6 +201,7 @@ impl Default for RailLimits { max_connections_total: 32, max_inflight_bytes_per_rail: 8 * 1024 * 1024 * 1024u64, max_inflight_bytes_total: 32 * 1024 * 1024 * 1024u64, + task_max_stripes: 0, io_timeout: Duration::from_secs(30), rail_cooldown: DEFAULT_RAIL_COOLDOWN, } @@ -740,10 +758,20 @@ fn pick_rail( /// then distribute each endpoint's stripes over the rails allowed to reach /// that endpoint, according to the policy, keeping one task per (rail, /// endpoint) pair so the byte count per task can be verified afterwards. +/// Total bytes of one stripe batch according to the validated placement. +fn stripe_batch_bytes(stripes: &[u32], placement: &ValidatedPlacement) -> u64 { + stripes + .iter() + .filter_map(|stripe| placement.stripes.get(stripe)) + .map(|info| info.length) + .sum() +} + fn build_plan( placement: &ValidatedPlacement, rails: &[Arc], policy: RailSelectPolicy, + limits: &RailLimits, ) -> Result { // Group stripes by endpoint first. let mut by_endpoint: BTreeMap, Vec<(u32, u64)>> = BTreeMap::new(); @@ -792,18 +820,27 @@ fn build_plan( bytes_per_rail[rail] += length; assigned[rail] += length; } - for (rail_index, (stripe_list, bytes)) in + for (rail_index, (stripe_list, _bytes)) in per_rail.into_iter().zip(bytes_per_rail).enumerate() { if stripe_list.is_empty() { continue; } - tasks.push(TaskSpec { - rail_index, - endpoint: Arc::clone(&endpoint), - stripes: stripe_list, - bytes, - }); + // Split into per-connection batches for intra-rail + // concurrency (bounded queue depth per rail). + let batch = if limits.task_max_stripes == 0 { + stripe_list.len() + } else { + limits.task_max_stripes.max(1) + }; + for chunk in stripe_list.chunks(batch) { + tasks.push(TaskSpec { + rail_index, + endpoint: Arc::clone(&endpoint), + stripes: chunk.to_vec(), + bytes: stripe_batch_bytes(chunk, placement), + }); + } } } // Pin the whole endpoint to one allowed rail. @@ -1211,7 +1248,7 @@ impl MultiRailClient { .map(|rail| (rail.index, Arc::clone(rail))) .collect(); - let plan = build_plan(&placement, &rails, self.policy)?; + let plan = build_plan(&placement, &rails, self.policy, &self.limits)?; let waves = plan_waves(plan.tasks, &self.limits); // Global deadline: every wave gets a full io_timeout, plus one extra @@ -1671,11 +1708,12 @@ mod tests { #[test] fn rail_config_parses_optional_fields() { assert_eq!(RailConfig::parse("mlx5_0").unwrap().port, 1); - let full = RailConfig::parse("irdma0:1:5:2").unwrap(); + let full = RailConfig::parse("irdma0:1:5:2:4096").unwrap(); assert_eq!(full.device, "irdma0"); assert_eq!(full.port, 1); assert_eq!(full.gid_index, 5); assert_eq!(full.weight, 2); + assert_eq!(full.mtu, 4096); assert!(RailConfig::parse(" ").is_none()); } @@ -1747,7 +1785,7 @@ mod tests { let chunks: Vec<_> = (0..8).map(|i| chunk(i, "10.0.0.1", 8, "")).collect(); let placement = validate_placement(&desc, &chunks).unwrap(); let rails = test_rails(2); - let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded).unwrap(); + let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded, &RailLimits::default()).unwrap(); assert_eq!(plan.task_count(), 2); let total: u64 = plan.tasks.iter().map(|t| t.bytes).sum(); assert_eq!(total, 64); @@ -1781,7 +1819,7 @@ mod tests { unhealthy_until: Mutex::new(None), gid_v4: None, }); - let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded).unwrap(); + let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded, &RailLimits::default()).unwrap(); let rail0: u64 = plan .tasks .iter() @@ -1814,7 +1852,7 @@ mod tests { gid_v4: None, }); let plan = - build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded).expect("plan builds"); + build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded, &RailLimits::default()).expect("plan builds"); // dev0 pinned to 10.0.0.2; 10.0.0.1 falls to dev1. assert_eq!(plan.tasks.len(), 2); for task in &plan.tasks { @@ -1832,7 +1870,7 @@ mod tests { unhealthy_until: Mutex::new(None), gid_v4: None, }); - assert!(build_plan(&placement, &strict, RailSelectPolicy::LeastLoaded).is_err()); + assert!(build_plan(&placement, &strict, RailSelectPolicy::LeastLoaded, &RailLimits::default()).is_err()); } #[test] @@ -1841,7 +1879,7 @@ mod tests { let chunks: Vec<_> = (0..8).map(|i| chunk(i, "10.0.0.1", 8, "")).collect(); let placement = validate_placement(&desc, &chunks).unwrap(); let rails = test_rails(2); - let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity).unwrap(); + let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity, &RailLimits::default()).unwrap(); assert_eq!(plan.task_count(), 1); assert_eq!(plan.tasks[0].bytes, 64); assert_eq!(plan.tasks[0].stripes.len(), 8); @@ -1862,7 +1900,7 @@ mod tests { unhealthy_until: Mutex::new(None), gid_v4: Some("10.0.0.5".parse().unwrap()), }); - let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity).unwrap(); + let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity, &RailLimits::default()).unwrap(); assert_eq!(plan.tasks.len(), 1); assert_eq!(plan.tasks[0].rail_index, 1); } diff --git a/kv-service/client-rs/src/rdma.rs b/kv-service/client-rs/src/rdma.rs index 964e5ae..22a6841 100644 --- a/kv-service/client-rs/src/rdma.rs +++ b/kv-service/client-rs/src/rdma.rs @@ -65,6 +65,9 @@ pub struct RdmaClientConfig { pub io_timeout: Option, /// Bound on the TCP connect phase itself. pub connect_timeout: Option, + /// RC path MTU in bytes (512/1024/2048/4096). Must not exceed the network + /// MTU; jumbo-frame RoCE fabrics want 4096. + pub path_mtu: u16, } impl RdmaClientConfig { @@ -77,6 +80,7 @@ impl RdmaClientConfig { gid_index: 3, io_timeout: None, connect_timeout: None, + path_mtu: 1024, } } @@ -103,6 +107,22 @@ impl RdmaClientConfig { self.connect_timeout = Some(timeout); self } + + /// Set the RC path MTU in bytes (rounded down to 512/1024/2048/4096). + pub fn with_path_mtu(mut self, bytes: u16) -> Self { + self.path_mtu = bytes; + self + } +} + +/// Map a byte count onto the closest supported RC path MTU. +fn path_mtu_enum(bytes: u16) -> ibv_mtu::Type { + match bytes { + 0..=512 => ibv_mtu::IBV_MTU_512, + 513..=1024 => ibv_mtu::IBV_MTU_1024, + 1025..=2048 => ibv_mtu::IBV_MTU_2048, + _ => ibv_mtu::IBV_MTU_4096, + } } /// Outcome of a descriptor GET: how many bytes the server placed in the @@ -411,7 +431,7 @@ impl RdmaClient { }; write_hello(&mut stream, local)?; let remote = read_hello(&mut stream)?; - transition_qp_to_rtr(qp, &remote, config.port, config.gid_index)?; + transition_qp_to_rtr(qp, &remote, config.port, config.gid_index, config.path_mtu)?; transition_qp_to_rts(qp, local.psn)?; Ok(stream) })(); @@ -1103,11 +1123,12 @@ fn transition_qp_to_rtr( remote: &QpInfo, port: u8, gid_index: u8, + path_mtu_bytes: u16, ) -> Result<()> { unsafe { let mut attr: ibv_qp_attr = std::mem::zeroed(); attr.qp_state = ibv_qp_state::IBV_QPS_RTR; - attr.path_mtu = ibv_mtu::IBV_MTU_1024; + attr.path_mtu = path_mtu_enum(path_mtu_bytes); attr.dest_qp_num = remote.qpn; attr.rq_psn = remote.psn; attr.max_dest_rd_atomic = 1; diff --git a/kv-service/client-rs/tests/multirail_e2e.rs b/kv-service/client-rs/tests/multirail_e2e.rs index b2408e4..a1985b9 100644 --- a/kv-service/client-rs/tests/multirail_e2e.rs +++ b/kv-service/client-rs/tests/multirail_e2e.rs @@ -31,6 +31,11 @@ fn rails_from_env() -> Vec { if let Some(endpoints) = pins.get(&rail.device) { rail.endpoints = endpoints.clone(); } + if let Ok(mtu) = std::env::var("CS_MR_RAIL_MTU") { + if let Ok(mtu) = mtu.trim().parse::() { + rail.mtu = mtu; + } + } rail }) .collect() @@ -162,9 +167,16 @@ fn seed(namespace: &str, key: &str, size_mb: usize) -> Fixture { } fn limits() -> RailLimits { + // CS_MR_TASK_MAX_STRIPES > 0 exercises the intra-rail task split (several + // connections per rail) in addition to the plain per-endpoint layout. + let task_max_stripes = std::env::var("CS_MR_TASK_MAX_STRIPES") + .ok() + .and_then(|value| value.trim().parse::().ok()) + .unwrap_or(0); RailLimits { io_timeout: Duration::from_secs(20), rail_cooldown: Duration::from_millis(500), + task_max_stripes, ..RailLimits::default() } } diff --git a/kv-service/server/src/rdma/qp.rs b/kv-service/server/src/rdma/qp.rs index 6f758af..29ee19a 100644 --- a/kv-service/server/src/rdma/qp.rs +++ b/kv-service/server/src/rdma/qp.rs @@ -134,11 +134,17 @@ impl RcQp { } /// Transition to RTR (Ready-To-Receive). Requires remote QP info + path MTU + GID index. - pub fn to_rtr(&self, remote: &QpInfo, port_num: u8, gid_index: u8) -> Result<()> { + pub fn to_rtr( + &self, + remote: &QpInfo, + port_num: u8, + gid_index: u8, + path_mtu: ibv_mtu::Type, + ) -> Result<()> { unsafe { let mut attr: ibv_qp_attr = std::mem::zeroed(); attr.qp_state = ibv_qp_state::IBV_QPS_RTR; - attr.path_mtu = ibv_mtu::IBV_MTU_1024; // align with hardware active_mtu + attr.path_mtu = path_mtu; // align with hardware/network MTU attr.dest_qp_num = remote.qpn; attr.rq_psn = remote.psn; attr.max_dest_rd_atomic = 1; diff --git a/kv-service/server/src/rdma/server.rs b/kv-service/server/src/rdma/server.rs index ef6801b..eb0081a 100644 --- a/kv-service/server/src/rdma/server.rs +++ b/kv-service/server/src/rdma/server.rs @@ -30,7 +30,7 @@ use crate::rdma::wire::{ use crate::router::ObjectKey; use crate::KVServiceContext; use anyhow::{anyhow, Result}; -use rdma_sys::ibv_access_flags; +use rdma_sys::{ibv_access_flags, ibv_mtu}; use std::net::{TcpListener, TcpStream}; use std::ptr::NonNull; use std::sync::Arc; @@ -265,7 +265,7 @@ fn handle_client( let remote = wire::recv_hello(&mut stream)?; wire::send_hello(&mut stream, &qp.local)?; - qp.to_rtr(&remote, port_num, gid_index)?; + qp.to_rtr(&remote, port_num, gid_index, path_mtu_from_env())?; qp.to_rts()?; tracing::info!( "RDMA QP established: local_qpn={} remote_qpn={}", @@ -974,6 +974,21 @@ fn descriptor_meta_from_req( Ok(meta) } +/// RC path MTU override (bytes) for listener QPs: CS_RDMA_PATH_MTU= +/// 512|1024|2048|4096, default 1024 (safe for standard 1500-byte networks). +fn path_mtu_from_env() -> ibv_mtu::Type { + match std::env::var("CS_RDMA_PATH_MTU") + .ok() + .and_then(|value| value.trim().parse::().ok()) + { + Some(mtu) if mtu <= 512 => ibv_mtu::IBV_MTU_512, + Some(mtu) if mtu <= 1024 => ibv_mtu::IBV_MTU_1024, + Some(mtu) if mtu <= 2048 => ibv_mtu::IBV_MTU_2048, + Some(mtu) if mtu <= 4096 => ibv_mtu::IBV_MTU_4096, + _ => ibv_mtu::IBV_MTU_1024, + } +} + fn chunk_is_local(ctx: &KVServiceContext, location: &crate::metadata::ChunkLocation) -> bool { let local_node_id = std::env::var("CS_NODE_ID") .ok() From 642a468951aac409df4153d975b0a1a1f029bf22 Mon Sep 17 00:00:00 2001 From: wzu Date: Sat, 26 Sep 2026 12:16:41 +0800 Subject: [PATCH 3/7] multirail: Python Worker FFI (cs_mr_*), cooperative cancel, rail endpoint pin syntax - rdma-ffi exposes MultiRailClient over a C ABI: cs_mr_new/read/ rail_stats/free with CsMrDescriptor/CsMrChunk mirroring the gRPC fields; Python ctypes binding MultiRailReader in contextstore.storage.multirail_client closes the KVConnector Worker loop (verified end to end: 128MiB over 2 rails, byte-exact balance). - CancelToken: shared-flag cooperative cancellation checked at wave boundaries, backpressure waits and reply collection; cancelled reads still drain in-flight replies and quiesce connections before returning (memory-safety contract unchanged). - RailConfig::parse gains an @host[:port][;...] endpoint whitelist suffix for rail-optimized pinning from a plain spec string. - Python integration suite (pytest tests/) passes 5/5. --- Cargo.lock | 1 + .../client-rs/src/bin/multirail_bench.rs | 29 ++ kv-service/client-rs/src/multirail.rs | 129 ++++++- kv-service/client-rs/src/rdma.rs | 21 +- kv-service/rdma-ffi/Cargo.toml | 1 + kv-service/rdma-ffi/src/lib.rs | 2 + kv-service/rdma-ffi/src/multirail.rs | 320 ++++++++++++++++++ kv-service/server/src/rdma/qp.rs | 3 +- kv-service/server/src/rdma/server.rs | 16 +- src/contextstore/storage/multirail_client.py | 269 +++++++++++++++ 10 files changed, 775 insertions(+), 16 deletions(-) create mode 100644 kv-service/rdma-ffi/src/multirail.rs create mode 100644 src/contextstore/storage/multirail_client.py diff --git a/Cargo.lock b/Cargo.lock index d063aef..2ed47fa 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -363,6 +363,7 @@ name = "contextstore-rdma-ffi" version = "0.1.0" dependencies = [ "anyhow", + "contextstore-client-rs", "libc", "rdma-sys", ] diff --git a/kv-service/client-rs/src/bin/multirail_bench.rs b/kv-service/client-rs/src/bin/multirail_bench.rs index 4dcb29c..37509a1 100644 --- a/kv-service/client-rs/src/bin/multirail_bench.rs +++ b/kv-service/client-rs/src/bin/multirail_bench.rs @@ -61,6 +61,14 @@ struct Args { /// (rail, endpoint). #[arg(long, default_value = "0")] task_max_stripes: usize, + /// GRH hop limit for all rails (routed RoCE needs more than 1). + #[arg(long, default_value = "1")] + hop_limit: u8, + /// Endpoint rewrite map for reachability indirection, e.g. + /// `10.0.0.2:50053=127.0.0.1:15053`: placement endpoints on the left are + /// dialed via the right (RDMA GIDs are exchanged in-band and unaffected). + #[arg(long, value_delimiter = ',')] + endpoint_map: Vec, #[arg(long, default_value = "5")] iters: usize, /// Destination buffer size in MiB (>= object size). @@ -303,6 +311,24 @@ fn main() -> Result<()> { return Err(anyhow!("lookup returned no placement (is the object striped?)")); } let original_lookup = lookup.clone(); + // Reachability indirection: rewrite placement endpoints through the map + // (the RDMA GID exchange travels inside the control stream, so the data + // path is unaffected by the TCP detour). + let endpoint_map: HashMap = args + .endpoint_map + .iter() + .filter_map(|spec| spec.split_once('=')) + .map(|(from, to)| (from.trim().to_string(), to.trim().to_string())) + .collect(); + if !endpoint_map.is_empty() { + if let Some(placement) = lookup.placement.as_mut() { + for chunk in placement.chunks.iter_mut() { + if let Some(to) = endpoint_map.get(&chunk.rdma_endpoint) { + chunk.rdma_endpoint = to.clone(); + } + } + } + } if !args.alternate_endpoints.is_empty() { // Spread the node's stripes over its listeners so several rails can // carry the object concurrently. @@ -327,6 +353,9 @@ fn main() -> Result<()> { if args.qp_mtu != 1024 { rail.mtu = args.qp_mtu; } + if args.hop_limit != 1 { + rail.hop_limit = args.hop_limit; + } rail }) }) diff --git a/kv-service/client-rs/src/multirail.rs b/kv-service/client-rs/src/multirail.rs index 0143ad4..c28c284 100644 --- a/kv-service/client-rs/src/multirail.rs +++ b/kv-service/client-rs/src/multirail.rs @@ -92,6 +92,8 @@ pub struct RailConfig { /// RC path MTU in bytes; 4096 on jumbo-frame fabrics, 1024 (default) on /// standard 1500-byte networks. pub mtu: u16, + /// GRH hop limit; routed RoCE fabrics need more than the default 1. + pub hop_limit: u8, /// Optional endpoint whitelist (`host:port` or bare `host`). Empty means /// the rail may serve any endpoint. Rail-optimized fabrics pin rail k to /// fabric k; soft-RoCE cross-wired testbeds use it the same way. @@ -106,6 +108,7 @@ impl RailConfig { gid_index: 3, weight: 1, mtu: 1024, + hop_limit: 1, endpoints: Vec::new(), } } @@ -131,6 +134,12 @@ impl RailConfig { self } + /// Set the GRH hop limit (routed RoCE needs more than 1). + pub fn with_hop_limit(mut self, hops: u8) -> Self { + self.hop_limit = hops; + self + } + /// Restrict this rail to the given endpoints (`host:port` or bare host). pub fn with_endpoints(mut self, endpoints: Vec) -> Self { self.endpoints = endpoints; @@ -149,8 +158,21 @@ impl RailConfig { .any(|allowed| allowed == endpoint || *allowed == host) } - /// Parse `device[:port[:gid[:weight[:mtu]]]]` (omitted parts keep defaults). + /// Parse `device[:port[:gid[:weight[:mtu]]]][@ep[;ep...]]` (omitted + /// parts keep defaults). The optional `@` suffix pins the rail to an + /// endpoint whitelist (`host:port` or bare host, `;`-separated). pub fn parse(spec: &str) -> Option { + let (spec, endpoints) = match spec.split_once('@') { + Some((left, right)) => ( + left, + right + .split(';') + .map(|ep| ep.trim().to_string()) + .filter(|ep| !ep.is_empty()) + .collect::>(), + ), + None => (spec, Vec::new()), + }; let mut parts = spec.split(':'); let device = parts.next()?.trim(); if device.is_empty() { @@ -169,6 +191,9 @@ impl RailConfig { if let Some(mtu) = parts.next().and_then(|m| m.trim().parse().ok()) { config = config.with_mtu(mtu); } + if !endpoints.is_empty() { + config = config.with_endpoints(endpoints); + } Some(config) } } @@ -279,6 +304,7 @@ pub enum MultiRailError { actual: String, }, WorkerPanic { rail: String, endpoint: String }, + Cancelled, } impl fmt::Display for MultiRailError { @@ -344,12 +370,50 @@ impl fmt::Display for MultiRailError { Self::WorkerPanic { rail, endpoint } => { write!(f, "rail {rail} worker for {endpoint} panicked") } + Self::Cancelled => { + write!(f, "read cancelled by caller") + } } } } impl std::error::Error for MultiRailError {} +/// Cooperative cancellation for multi-rail reads. +/// +/// Clones share the same flag. Cancellation is observed at wave boundaries, +/// during backpressure waits, and between reply-collection cycles; the +/// in-flight tasks of the current wave are still drained (bounded by the +/// read deadline) and every participating connection is quiesced before the +/// read returns, so cancelling never violates the memory-safety contract. +#[derive(Clone, Default)] +pub struct CancelToken { + flag: Arc, +} + +impl CancelToken { + pub fn new() -> Self { + Self::default() + } + + /// Request cancellation. + pub fn cancel(&self) { + self.flag.store(true, Ordering::Release); + } + + pub fn is_cancelled(&self) -> bool { + self.flag.load(Ordering::Acquire) + } + + fn check(&self) -> Result<(), MultiRailError> { + if self.is_cancelled() { + Err(MultiRailError::Cancelled) + } else { + Ok(()) + } + } +} + // --------------------------------------------------------------------------- // Topology // --------------------------------------------------------------------------- @@ -1175,7 +1239,19 @@ impl MultiRailClient { buffer: &mut [u8], ) -> Result { let base = buffer.as_mut_ptr() as usize; - self.read_impl(descriptor, chunks, base, buffer.len(), false) + self.read_impl(descriptor, chunks, base, buffer.len(), false, None) + } + + /// [`Self::read_object_into`] with cooperative cancellation. + pub fn read_object_into_cancelled( + &self, + descriptor: &pb::ObjectDescriptor, + chunks: &[pb::PlacementChunk], + buffer: &mut [u8], + cancel: &CancelToken, + ) -> Result { + let base = buffer.as_mut_ptr() as usize; + self.read_impl(descriptor, chunks, base, buffer.len(), false, Some(cancel)) } /// Convenience wrapper taking a gRPC [`crate::ObjectLookup`] result. @@ -1214,7 +1290,7 @@ impl MultiRailClient { have: len, }); } - self.read_impl(descriptor, chunks, ptr as usize, len, sticky_registration) + self.read_impl(descriptor, chunks, ptr as usize, len, sticky_registration, None) } fn read_impl( @@ -1224,6 +1300,7 @@ impl MultiRailClient { base: usize, len: usize, sticky: bool, + cancel: Option<&CancelToken>, ) -> Result { let placement = validate_placement(descriptor, chunks)?; if placement.object_size > len as u64 { @@ -1268,6 +1345,9 @@ impl MultiRailClient { let mut verified_bytes = 0u64; 'waves: for wave in waves { + if let Some(cancel) = cancel { + cancel.check()?; + } // ---- dispatch this wave ---- let mut dispatched = 0usize; for task in &wave { @@ -1275,7 +1355,7 @@ impl MultiRailClient { continue; }; // Backpressure: wait for per-rail and total byte headroom. - if !self.await_headroom(rail, task.bytes, deadline) { + if !self.await_headroom(rail, task.bytes, deadline, cancel) { failure = Some(MultiRailError::Timeout { rail: rail.config.device.clone(), endpoint: task.endpoint.to_string(), @@ -1329,15 +1409,29 @@ impl MultiRailClient { }); break; } - match reply_rx.recv_timeout(remaining) { + if let Some(cancel) = cancel { + if let Err(error) = cancel.check() { + // Drain with a short grace period, then quiesce via + // the failure path below (join still guarantees that + // no work request outlives this call). + failure = Some(error); + break; + } + } + // Bound each wait so cancellation is observed promptly. + let slice = remaining.min(Duration::from_millis(100)); + match reply_rx.recv_timeout(slice) { Ok(reply) => replies.push(reply), Err(mpsc::RecvTimeoutError::Timeout) => { - failure = Some(MultiRailError::Timeout { - rail: "any".into(), - endpoint: "any".into(), - after_ms: self.limits.io_timeout.as_millis(), - }); - break; + if deadline <= Instant::now() { + failure = Some(MultiRailError::Timeout { + rail: "any".into(), + endpoint: "any".into(), + after_ms: self.limits.io_timeout.as_millis(), + }); + break; + } + continue; } Err(mpsc::RecvTimeoutError::Disconnected) => { failure = Some(MultiRailError::WorkerPanic { @@ -1498,8 +1592,19 @@ impl MultiRailClient { /// Wait until `bytes` more in-flight traffic fits the per-rail and total /// budgets, or the deadline passes. - fn await_headroom(&self, rail: &Rail, bytes: u64, deadline: Instant) -> bool { + fn await_headroom( + &self, + rail: &Rail, + bytes: u64, + deadline: Instant, + cancel: Option<&CancelToken>, + ) -> bool { loop { + if let Some(cancel) = cancel { + if cancel.is_cancelled() { + return false; + } + } let rail_inflight = rail.stats.inflight_bytes.load(Ordering::Relaxed); if rail_inflight + bytes <= self.limits.max_inflight_bytes_per_rail { let total: u64 = self diff --git a/kv-service/client-rs/src/rdma.rs b/kv-service/client-rs/src/rdma.rs index 22a6841..56a3fdd 100644 --- a/kv-service/client-rs/src/rdma.rs +++ b/kv-service/client-rs/src/rdma.rs @@ -68,6 +68,8 @@ pub struct RdmaClientConfig { /// RC path MTU in bytes (512/1024/2048/4096). Must not exceed the network /// MTU; jumbo-frame RoCE fabrics want 4096. pub path_mtu: u16, + /// GRH hop limit. 1 fits same-subnet fabrics; routed RoCE needs more. + pub hop_limit: u8, } impl RdmaClientConfig { @@ -81,6 +83,7 @@ impl RdmaClientConfig { io_timeout: None, connect_timeout: None, path_mtu: 1024, + hop_limit: 1, } } @@ -113,6 +116,12 @@ impl RdmaClientConfig { self.path_mtu = bytes; self } + + /// Set the GRH hop limit (routed RoCE fabrics need more than 1). + pub fn with_hop_limit(mut self, hops: u8) -> Self { + self.hop_limit = hops; + self + } } /// Map a byte count onto the closest supported RC path MTU. @@ -431,7 +440,14 @@ impl RdmaClient { }; write_hello(&mut stream, local)?; let remote = read_hello(&mut stream)?; - transition_qp_to_rtr(qp, &remote, config.port, config.gid_index, config.path_mtu)?; + transition_qp_to_rtr( + qp, + &remote, + config.port, + config.gid_index, + config.path_mtu, + config.hop_limit, + )?; transition_qp_to_rts(qp, local.psn)?; Ok(stream) })(); @@ -1124,6 +1140,7 @@ fn transition_qp_to_rtr( port: u8, gid_index: u8, path_mtu_bytes: u16, + hop_limit: u8, ) -> Result<()> { unsafe { let mut attr: ibv_qp_attr = std::mem::zeroed(); @@ -1136,7 +1153,7 @@ fn transition_qp_to_rtr( attr.ah_attr.is_global = 1; attr.ah_attr.port_num = port; attr.ah_attr.grh.dgid = remote.gid; - attr.ah_attr.grh.hop_limit = 1; + attr.ah_attr.grh.hop_limit = hop_limit; attr.ah_attr.grh.sgid_index = gid_index; let mask = ibv_qp_attr_mask::IBV_QP_STATE | ibv_qp_attr_mask::IBV_QP_AV diff --git a/kv-service/rdma-ffi/Cargo.toml b/kv-service/rdma-ffi/Cargo.toml index 499813c..2cb7b7a 100644 --- a/kv-service/rdma-ffi/Cargo.toml +++ b/kv-service/rdma-ffi/Cargo.toml @@ -14,3 +14,4 @@ path = "src/lib.rs" rdma-sys = "0.3" anyhow = "1" libc = "0.2" +contextstore-client-rs = { path = "../client-rs", features = ["rdma"] } diff --git a/kv-service/rdma-ffi/src/lib.rs b/kv-service/rdma-ffi/src/lib.rs index 843e73f..f8be2a9 100644 --- a/kv-service/rdma-ffi/src/lib.rs +++ b/kv-service/rdma-ffi/src/lib.rs @@ -37,6 +37,8 @@ //! void cs_rdma_client_free(void* client); //! ``` +pub mod multirail; + use anyhow::{anyhow, Result}; use rdma_sys::*; use std::ffi::{c_char, c_int, CStr}; diff --git a/kv-service/rdma-ffi/src/multirail.rs b/kv-service/rdma-ffi/src/multirail.rs new file mode 100644 index 0000000..414fbe3 --- /dev/null +++ b/kv-service/rdma-ffi/src/multirail.rs @@ -0,0 +1,320 @@ +//! Multi-rail C ABI: exposes `contextstore_client_rs::multirail::MultiRailClient` +//! to Python ctypes, so a Python KVConnector Worker can read one striped +//! object in parallel over several local RDMA devices. +//! +//! ```c +//! // Create a reader over N rails. rail_specs[i] = "device[:port[:gid[:weight[:mtu]]]]". +//! // io_timeout_ms bounds each operation (and connect). NULL on error. +//! void* cs_mr_new(const char* const* rail_specs, uint32_t rail_count, uint64_t io_timeout_ms); +//! +//! // Object identity + placement, mirroring the gRPC ObjectDescriptor / +//! // PlacementChunk the Python side obtained from LookupObject. +//! typedef struct { +//! const char* namespace; +//! const char* object_key; +//! const char* object_handle; +//! uint64_t object_generation; +//! const char* content_etag; +//! uint64_t layout_version; +//! uint64_t size; +//! uint32_t is_striped; +//! uint32_t stripe_count; +//! uint64_t chunk_size; +//! } CsMrDescriptor; +//! typedef struct { +//! uint32_t stripe_index; +//! const char* rdma_endpoint; +//! uint64_t offset; +//! uint64_t length; +//! const char* checksum; // optional xxh3-64 lowercase hex; NULL/"" skips +//! } CsMrChunk; +//! +//! // Read the whole object into buffer (server RDMA-WRITEs each stripe at +//! // buffer + stripe_index * chunk_size). sticky=1 keeps per-rail registrations +//! // cached across calls — the buffer must then be a long-lived pinned pool +//! // region that outlives the reader and is never freed/reused otherwise. +//! // Returns bytes read (>=0) or -1; human-readable error text in err_buf. +//! int64_t cs_mr_read(void* reader, +//! const CsMrDescriptor* descriptor, +//! const CsMrChunk* chunks, uint32_t chunk_count, +//! uint8_t* buffer, uint64_t buffer_len, +//! int32_t sticky, +//! char* err_buf, uint32_t err_buf_len); +//! +//! // Per-rail stats snapshot; returns the number of rails written to `out`. +//! typedef struct { +//! uint32_t index, healthy, cooldown_ms; +//! uint64_t requests_ok, requests_err, bytes_read, timeouts; +//! uint64_t connections_created, connections_quiesced; +//! uint64_t inflight_requests, inflight_bytes; +//! char device[64]; +//! char topology[160]; +//! } CsMrRailStats; +//! int32_t cs_mr_rail_stats(void* reader, CsMrRailStats* out, uint32_t max); +//! +//! void cs_mr_free(void* reader); +//! ``` + +use contextstore_client_rs::multirail::{ + MultiRailClient, MultiRailError, RailConfig, RailLimits, RailSelectPolicy, +}; +use contextstore_client_rs::pb; +use std::ffi::{c_char, CStr}; +use std::os::raw::{c_int, c_void}; +use std::time::Duration; + +/// Opaque reader handle body. +pub struct MrReader { + client: MultiRailClient, +} + +#[repr(C)] +pub struct CsMrDescriptor { + pub namespace: *const c_char, + pub object_key: *const c_char, + pub object_handle: *const c_char, + pub object_generation: u64, + pub content_etag: *const c_char, + pub layout_version: u64, + pub size: u64, + pub is_striped: u32, + pub stripe_count: u32, + pub chunk_size: u64, +} + +#[repr(C)] +pub struct CsMrChunk { + pub stripe_index: u32, + pub rdma_endpoint: *const c_char, + pub offset: u64, + pub length: u64, + pub checksum: *const c_char, +} + +#[repr(C)] +pub struct CsMrRailStats { + pub index: u32, + pub healthy: u32, + pub cooldown_ms: u32, + pub requests_ok: u64, + pub requests_err: u64, + pub bytes_read: u64, + pub timeouts: u64, + pub connections_created: u64, + pub connections_quiesced: u64, + pub inflight_requests: u64, + pub inflight_bytes: u64, + pub device: [c_char; 64], + pub topology: [c_char; 160], +} + +fn write_err(err_buf: *mut c_char, err_buf_len: u32, message: &str) { + if err_buf.is_null() || err_buf_len == 0 { + return; + } + let bytes = message.as_bytes(); + let capacity = err_buf_len as usize - 1; + let len = bytes.len().min(capacity); + unsafe { + std::ptr::copy_nonoverlapping(bytes.as_ptr(), err_buf as *mut u8, len); + *err_buf.add(len) = 0; + } +} + +unsafe fn cstr<'a>(ptr: *const c_char) -> Result<&'a str, MultiRailError> { + if ptr.is_null() { + return Err(MultiRailError::InvalidPlacement("null string".into())); + } + CStr::from_ptr(ptr).to_str().map_err(|_| { + MultiRailError::InvalidPlacement("non-utf8 string in descriptor".into()) + }) +} + +#[no_mangle] +pub unsafe extern "C" fn cs_mr_new( + rail_specs: *const *const c_char, + rail_count: u32, + io_timeout_ms: u64, +) -> *mut c_void { + let mut rails = Vec::with_capacity(rail_count as usize); + for index in 0..rail_count as usize { + let spec = match rail_specs.add(index).read().as_ref().and_then(|p| CStr::from_ptr(p).to_str().ok()) { + Some(spec) => spec, + None => return std::ptr::null_mut(), + }; + match RailConfig::parse(spec) { + Some(config) => rails.push(config), + None => return std::ptr::null_mut(), + } + } + let limits = RailLimits { + io_timeout: Duration::from_millis(io_timeout_ms.max(1)), + ..RailLimits::default() + }; + match MultiRailClient::new(rails) { + Ok(client) => Box::into_raw(Box::new(MrReader { + client: client.with_limits(limits).with_policy(RailSelectPolicy::LeastLoaded), + })) as *mut c_void, + Err(_) => std::ptr::null_mut(), + } +} + +fn to_placement( + descriptor: &CsMrDescriptor, + chunks: &[CsMrChunk], +) -> Result<(pb::ObjectDescriptor, Vec), MultiRailError> { + unsafe { + let descriptor = pb::ObjectDescriptor { + key: Some(pb::ObjectKey { + namespace: cstr(descriptor.namespace)?.to_string(), + object_key: cstr(descriptor.object_key)?.to_string(), + }), + object_handle: cstr(descriptor.object_handle)?.to_string(), + object_generation: descriptor.object_generation, + content_etag: cstr(descriptor.content_etag)?.to_string(), + layout_version: descriptor.layout_version, + size: descriptor.size, + is_striped: descriptor.is_striped != 0, + stripe_count: descriptor.stripe_count, + chunk_size: descriptor.chunk_size, + }; + let mut out = Vec::with_capacity(chunks.len()); + for chunk in chunks { + out.push(pb::PlacementChunk { + stripe_index: chunk.stripe_index, + node_id: String::new(), + grpc_endpoint: String::new(), + rdma_endpoint: cstr(chunk.rdma_endpoint)?.to_string(), + device_id: 0, + storage_handle: String::new(), + offset: chunk.offset, + length: chunk.length, + checksum: if chunk.checksum.is_null() { + String::new() + } else { + cstr(chunk.checksum)?.to_string() + }, + }); + } + Ok((descriptor, out)) + } +} + +#[no_mangle] +pub unsafe extern "C" fn cs_mr_read( + reader: *mut c_void, + descriptor: *const CsMrDescriptor, + chunks: *const CsMrChunk, + chunk_count: u32, + buffer: *mut u8, + buffer_len: u64, + sticky: c_int, + err_buf: *mut c_char, + err_buf_len: u32, +) -> i64 { + let reader = match (reader as *mut MrReader).as_mut() { + Some(reader) => reader, + None => { + write_err(err_buf, err_buf_len, "cs_mr_read: null reader"); + return -1; + } + }; + let descriptor = match descriptor.as_ref() { + Some(descriptor) => descriptor, + None => { + write_err(err_buf, err_buf_len, "cs_mr_read: null descriptor"); + return -1; + } + }; + if chunks.is_null() || chunk_count == 0 || buffer.is_null() { + write_err(err_buf, err_buf_len, "cs_mr_read: null chunks/buffer"); + return -1; + } + let chunk_slice = std::slice::from_raw_parts(chunks, chunk_count as usize); + let (pb_descriptor, pb_chunks) = match to_placement(descriptor, chunk_slice) { + Ok(value) => value, + Err(error) => { + write_err(err_buf, err_buf_len, &error.to_string()); + return -1; + } + }; + let result = if sticky != 0 { + reader.client.read_object_into_raw( + &pb_descriptor, + &pb_chunks, + buffer, + buffer_len as usize, + true, + ) + } else { + // SAFETY: the caller guarantees buffer..buffer+len stays valid and + // unmoved for the duration of the call (it is synchronous). + reader + .client + .read_object_into_raw(&pb_descriptor, &pb_chunks, buffer, buffer_len as usize, false) + }; + match result { + Ok(bytes) => bytes as i64, + Err(error) => { + write_err(err_buf, err_buf_len, &error.to_string()); + -1 + } + } +} + +/// Copy a Rust string into a fixed-size C buffer (truncating, always +/// NUL-terminated). +fn fill_str_field(dst: &mut [c_char], value: &str) { + dst[0] = 0; + let bytes = value.as_bytes(); + let len = bytes.len().min(dst.len() - 1); + for (index, byte) in bytes[..len].iter().enumerate() { + dst[index] = *byte as c_char; + } + dst[len] = 0; +} + +#[no_mangle] +pub unsafe extern "C" fn cs_mr_rail_stats( + reader: *mut c_void, + out: *mut CsMrRailStats, + max: u32, +) -> i32 { + let reader = match (reader as *mut MrReader).as_ref() { + Some(reader) => reader, + None => return -1, + }; + if out.is_null() || max == 0 { + return -1; + } + let snapshots = reader.client.rails_snapshot(); + let count = snapshots.len().min(max as usize); + for (index, snapshot) in snapshots[..count].iter().enumerate() { + let dst = out.add(index); + *dst = CsMrRailStats { + index: snapshot.index as u32, + healthy: u32::from(snapshot.healthy), + cooldown_ms: snapshot.cooldown_ms_remaining as u32, + requests_ok: snapshot.requests_ok, + requests_err: snapshot.requests_err, + bytes_read: snapshot.bytes_read, + timeouts: snapshot.timeouts, + connections_created: snapshot.connections_created, + connections_quiesced: snapshot.connections_quiesced, + inflight_requests: snapshot.inflight_requests, + inflight_bytes: snapshot.inflight_bytes, + device: [0; 64], + topology: [0; 160], + }; + fill_str_field(&mut (*dst).device, &snapshot.device); + fill_str_field(&mut (*dst).topology, &snapshot.topology.to_string()); + } + count as i32 +} + +#[no_mangle] +pub unsafe extern "C" fn cs_mr_free(reader: *mut c_void) { + if !reader.is_null() { + drop(Box::from_raw(reader as *mut MrReader)); + } +} diff --git a/kv-service/server/src/rdma/qp.rs b/kv-service/server/src/rdma/qp.rs index 29ee19a..1d39306 100644 --- a/kv-service/server/src/rdma/qp.rs +++ b/kv-service/server/src/rdma/qp.rs @@ -140,6 +140,7 @@ impl RcQp { port_num: u8, gid_index: u8, path_mtu: ibv_mtu::Type, + hop_limit: u8, ) -> Result<()> { unsafe { let mut attr: ibv_qp_attr = std::mem::zeroed(); @@ -158,7 +159,7 @@ impl RcQp { attr.ah_attr.port_num = port_num; attr.ah_attr.grh.dgid = remote.gid; attr.ah_attr.grh.flow_label = 0; - attr.ah_attr.grh.hop_limit = 1; + attr.ah_attr.grh.hop_limit = hop_limit; attr.ah_attr.grh.sgid_index = gid_index; attr.ah_attr.grh.traffic_class = 0; diff --git a/kv-service/server/src/rdma/server.rs b/kv-service/server/src/rdma/server.rs index eb0081a..8a99d70 100644 --- a/kv-service/server/src/rdma/server.rs +++ b/kv-service/server/src/rdma/server.rs @@ -265,7 +265,13 @@ fn handle_client( let remote = wire::recv_hello(&mut stream)?; wire::send_hello(&mut stream, &qp.local)?; - qp.to_rtr(&remote, port_num, gid_index, path_mtu_from_env())?; + qp.to_rtr( + &remote, + port_num, + gid_index, + path_mtu_from_env(), + hop_limit_from_env(), + )?; qp.to_rts()?; tracing::info!( "RDMA QP established: local_qpn={} remote_qpn={}", @@ -989,6 +995,14 @@ fn path_mtu_from_env() -> ibv_mtu::Type { } } +/// GRH hop limit override: CS_RDMA_HOP_LIMIT (default 1, same-subnet). +fn hop_limit_from_env() -> u8 { + std::env::var("CS_RDMA_HOP_LIMIT") + .ok() + .and_then(|value| value.trim().parse::().ok()) + .unwrap_or(1) +} + fn chunk_is_local(ctx: &KVServiceContext, location: &crate::metadata::ChunkLocation) -> bool { let local_node_id = std::env::var("CS_NODE_ID") .ok() diff --git a/src/contextstore/storage/multirail_client.py b/src/contextstore/storage/multirail_client.py new file mode 100644 index 0000000..8d5a2cf --- /dev/null +++ b/src/contextstore/storage/multirail_client.py @@ -0,0 +1,269 @@ +"""Multi-rail RDMA reader for ContextStore (Python ctypes binding). + +A single Worker reads one striped object in parallel over several local RDMA +devices ("rails"). The Rust side (:mod:`contextstore_rdma_ffi` ``cs_mr_*`` C +ABI) manages per-rail QPs/CQs/memory regions, health + cooldown, byte-exact +stripe balancing, backpressure, integrity verification (byte counts, chunk +counts, optional per-stripe xxh3-64), and late-write protection: a failed or +cancelled read quiesces every participating connection before returning. + +Typical use with the gRPC client:: + + from contextstore.kvservice_client.client import KVClient + from contextstore.storage.multirail_client import MultiRailReader + + grpc = KVClient("http://10.0.0.1:50051") + reader = MultiRailReader(["mlx5_0", "mlx5_1"]) # local HCAs + buffer = pinned_pool_region(size) # long-lived + + lookup = grpc.lookup_object("ns", "key") # descriptor+placement + n = reader.read_into(buffer, lookup, sticky=True) # all rails in parallel + for stats in reader.rail_stats(): + print(stats) # per-rail metrics + +``sticky=True`` keeps the per-rail registrations cached across reads; the +buffer must then be a long-lived pinned pool region that outlives the reader +and is never freed or reused for non-RDMA purposes while it is open. +""" + +from __future__ import annotations + +import ctypes +import os +from dataclasses import dataclass +from typing import Sequence + +from .rdma_client import _find_lib + +_ERR_LEN = 512 + + +class _CsMrDescriptor(ctypes.Structure): + _fields_ = [ + ("namespace", ctypes.c_char_p), + ("object_key", ctypes.c_char_p), + ("object_handle", ctypes.c_char_p), + ("object_generation", ctypes.c_uint64), + ("content_etag", ctypes.c_char_p), + ("layout_version", ctypes.c_uint64), + ("size", ctypes.c_uint64), + ("is_striped", ctypes.c_uint32), + ("stripe_count", ctypes.c_uint32), + ("chunk_size", ctypes.c_uint64), + ] + + +class _CsMrChunk(ctypes.Structure): + _fields_ = [ + ("stripe_index", ctypes.c_uint32), + ("rdma_endpoint", ctypes.c_char_p), + ("offset", ctypes.c_uint64), + ("length", ctypes.c_uint64), + ("checksum", ctypes.c_char_p), + ] + + +class _CsMrRailStats(ctypes.Structure): + _fields_ = [ + ("index", ctypes.c_uint32), + ("healthy", ctypes.c_uint32), + ("cooldown_ms", ctypes.c_uint32), + ("requests_ok", ctypes.c_uint64), + ("requests_err", ctypes.c_uint64), + ("bytes_read", ctypes.c_uint64), + ("timeouts", ctypes.c_uint64), + ("connections_created", ctypes.c_uint64), + ("connections_quiesced", ctypes.c_uint64), + ("inflight_requests", ctypes.c_uint64), + ("inflight_bytes", ctypes.c_uint64), + ("device", ctypes.c_char * 64), + ("topology", ctypes.c_char * 160), + ] + + +@dataclass +class RailStats: + index: int + device: str + topology: str + healthy: bool + cooldown_ms: int + requests_ok: int + requests_err: int + bytes_read: int + timeouts: int + connections_created: int + connections_quiesced: int + inflight_requests: int + inflight_bytes: int + + def __str__(self) -> str: + return ( + f"rail[{self.index}] {self.device} {self.topology} " + f"healthy={self.healthy} cooldown={self.cooldown_ms}ms " + f"ok={self.requests_ok} err={self.requests_err} tmo={self.timeouts} " + f"read={self.bytes_read >> 20}MiB conns=+{self.connections_created}" + f"/~{self.connections_quiesced} " + f"inflight=(req:{self.inflight_requests},B:{self.inflight_bytes})" + ) + + +class MultiRailError(RuntimeError): + """Raised when a multi-rail read fails; text comes from the typed Rust error.""" + + +class MultiRailReader: + """Read one striped object across several local RDMA devices.""" + + def __init__( + self, + rails: Sequence[str], + io_timeout_ms: int = 30_000, + lib_path: str | None = None, + ) -> None: + """``rails`` entries are ``device[:port[:gid[:weight[:mtu]]]]`` specs.""" + if not rails: + raise ValueError("at least one rail is required") + self._lib = ctypes.CDLL(lib_path or _find_lib(), use_errno=True) + self._setup_prototypes() + specs = [ctypes.c_char_p(r.encode()) for r in rails] + array = (ctypes.c_char_p * len(specs))(*specs) + self._handle = self._lib.cs_mr_new(array, len(specs), io_timeout_ms) + if not self._handle: + raise MultiRailError( + f"cs_mr_new failed (rails={list(rails)}); check device names/GIDs" + ) + + def _setup_prototypes(self) -> None: + lib = self._lib + lib.cs_mr_new.restype = ctypes.c_void_p + lib.cs_mr_new.argtypes = [ + ctypes.POINTER(ctypes.c_char_p), + ctypes.c_uint32, + ctypes.c_uint64, + ] + lib.cs_mr_read.restype = ctypes.c_int64 + lib.cs_mr_read.argtypes = [ + ctypes.c_void_p, + ctypes.POINTER(_CsMrDescriptor), + ctypes.POINTER(_CsMrChunk), + ctypes.c_uint32, + ctypes.c_void_p, + ctypes.c_uint64, + ctypes.c_int32, + ctypes.c_char_p, + ctypes.c_uint32, + ] + lib.cs_mr_rail_stats.restype = ctypes.c_int32 + lib.cs_mr_rail_stats.argtypes = [ + ctypes.c_void_p, + ctypes.POINTER(_CsMrRailStats), + ctypes.c_uint32, + ] + lib.cs_mr_free.argtypes = [ctypes.c_void_p] + lib.cs_mr_free.restype = None + + def read_into( + self, + buffer, # ctypes buffer / memoryview-compatible address + lookup, + sticky: bool = False, + buffer_addr: int | None = None, + buffer_len: int | None = None, + ) -> int: + """Read the looked-up object into ``buffer`` across all healthy rails. + + ``lookup`` is a gRPC ``LookupObjectResponse``-style object exposing + ``.descriptor`` (with namespace/object_key via ``.key``, handle, + generation, etag, layout, size, striping fields) and ``.placement`` + (with ``.chunks``). Returns the number of bytes placed in the buffer. + """ + descriptor = lookup.descriptor + placement = getattr(lookup, "placement", None) + if placement is None or not placement.chunks: + raise MultiRailError("lookup carried no placement chunks") + key = getattr(descriptor, "key", None) + if key is None: + raise MultiRailError("descriptor is missing its object key") + + addr = buffer_addr if buffer_addr is not None else ctypes.addressof(buffer) + length = buffer_len if buffer_len is not None else ctypes.sizeof(buffer) + + c_desc = _CsMrDescriptor( + namespace=key.namespace.encode(), + object_key=key.object_key.encode(), + object_handle=descriptor.object_handle.encode(), + object_generation=descriptor.object_generation, + content_etag=descriptor.content_etag.encode(), + layout_version=descriptor.layout_version, + size=descriptor.size, + is_striped=1 if descriptor.is_striped else 0, + stripe_count=descriptor.stripe_count, + chunk_size=descriptor.chunk_size, + ) + c_chunks = [ + _CsMrChunk( + stripe_index=chunk.stripe_index, + rdma_endpoint=chunk.rdma_endpoint.encode(), + offset=chunk.offset, + length=chunk.length, + checksum=(chunk.checksum or "").encode() or None, + ) + for chunk in placement.chunks + ] + array = (_CsMrChunk * len(c_chunks))(*c_chunks) + err = ctypes.create_string_buffer(_ERR_LEN) + result = self._lib.cs_mr_read( + self._handle, + ctypes.byref(c_desc), + array, + len(c_chunks), + ctypes.c_void_p(addr), + length, + 1 if sticky else 0, + err, + _ERR_LEN, + ) + if result < 0: + raise MultiRailError(err.value.decode(errors="replace") or "cs_mr_read failed") + return int(result) + + def rail_stats(self) -> list[RailStats]: + """Snapshot of per-rail health, throughput, error and in-flight stats.""" + max_rails = 16 + array = (_CsMrRailStats * max_rails)() + count = self._lib.cs_mr_rail_stats(self._handle, array, max_rails) + if count < 0: + raise MultiRailError("cs_mr_rail_stats failed") + stats = [] + for i in range(count): + entry = array[i] + stats.append( + RailStats( + index=entry.index, + device=entry.device.decode(errors="replace"), + topology=entry.topology.decode(errors="replace"), + healthy=bool(entry.healthy), + cooldown_ms=entry.cooldown_ms, + requests_ok=entry.requests_ok, + requests_err=entry.requests_err, + bytes_read=entry.bytes_read, + timeouts=entry.timeouts, + connections_created=entry.connections_created, + connections_quiesced=entry.connections_quiesced, + inflight_requests=entry.inflight_requests, + inflight_bytes=entry.inflight_bytes, + ) + ) + return stats + + def close(self) -> None: + if getattr(self, "_handle", None): + self._lib.cs_mr_free(self._handle) + self._handle = None + + def __del__(self): # noqa: D105 + try: + self.close() + except Exception: + pass From 0485f26d4f01da3e3a626fa4ef28e88485116081 Mon Sep 17 00:00:00 2001 From: wzu Date: Sat, 26 Sep 2026 13:43:27 +0800 Subject: [PATCH 4/7] multirail: per-rail latency and registered-memory metrics RailStats/RailSnapshot/FFI/Python now expose per-request latency (avg/max over successful reads) and the bytes currently pinned by cached registrations (accounted at register/evict time; the safe API returns to zero after each read, sticky pooling holds the buffer size). Closes the observability items of the requirements list. --- kv-service/client-rs/src/multirail.rs | 39 +++++++++++++++++++- kv-service/client-rs/src/rdma.rs | 15 +++++--- kv-service/rdma-ffi/src/multirail.rs | 6 +++ src/contextstore/storage/multirail_client.py | 13 ++++++- 4 files changed, 65 insertions(+), 8 deletions(-) diff --git a/kv-service/client-rs/src/multirail.rs b/kv-service/client-rs/src/multirail.rs index c28c284..3422a6b 100644 --- a/kv-service/client-rs/src/multirail.rs +++ b/kv-service/client-rs/src/multirail.rs @@ -482,6 +482,11 @@ struct RailStats { connections_quiesced: AtomicU64, inflight_requests: AtomicU64, inflight_bytes: AtomicU64, + /// Sum and max of per-request latencies (µs); avg = sum / requests_ok. + latency_us_sum: AtomicU64, + latency_us_max: AtomicU64, + /// Bytes currently pinned by cached memory registrations on this rail. + registered_bytes: AtomicU64, } /// Point-in-time view of one rail for observability. @@ -500,13 +505,16 @@ pub struct RailSnapshot { pub connections_quiesced: u64, pub inflight_requests: u64, pub inflight_bytes: u64, + pub latency_avg_us: u64, + pub latency_max_us: u64, + pub registered_bytes: u64, } impl fmt::Display for RailSnapshot { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!( f, - "rail[{}] {} {} healthy={} cooldown={}ms ok={} err={} tmo={} read={}MiB conns=+{}/~{} inflight=(req:{},B:{})", + "rail[{}] {} {} healthy={} cooldown={}ms ok={} err={} tmo={} read={}MiB conns=+{}/~{} inflight=(req:{},B:{}) lat=(avg:{}us,max:{}us) reg={}MiB", self.index, self.device, self.topology, @@ -520,6 +528,9 @@ impl fmt::Display for RailSnapshot { self.connections_quiesced, self.inflight_requests, self.inflight_bytes, + self.latency_avg_us, + self.latency_max_us, + self.registered_bytes >> 20, ) } } @@ -1093,21 +1104,37 @@ fn conn_worker( } => { let (expected_bytes, expected_chunks) = expected_task_outcome(&descriptor, &stripes); + let started = Instant::now(); let outcome = (|| -> Result, String> { // SAFETY: dst_base..dst_base+dst_len stays valid and // unmoved for the whole read call — the dispatcher joins // every worker before returning, and registrations for // non-sticky buffers are evicted synchronously at the end // of each read. + let before = client.registered_cache_bytes(); let view = unsafe { client .register_raw_buffer_cached(dst_base as *mut u8, dst_len) .map_err(|error| error.to_string())? }; + let registered_now = client.registered_cache_bytes(); + if registered_now > before { + rail.stats.registered_bytes.fetch_add( + (registered_now - before) as u64, + Ordering::Relaxed, + ); + } client .get_descriptor_stripes_into_view_detailed(&descriptor, &stripes, view, 0) .map_err(|error| error.to_string()) })(); + rail.stats + .latency_us_sum + .fetch_add(started.elapsed().as_micros() as u64, Ordering::Relaxed); + rail.stats.latency_us_max.fetch_max( + started.elapsed().as_micros() as u64, + Ordering::Relaxed, + ); let _ = reply.send(TaskReply { rail_index: rail.index, endpoint: Arc::clone(&endpoint), @@ -1117,7 +1144,8 @@ fn conn_worker( }); } Command::EvictRegistration { base, ack } => { - client.evict_registrations_for(base); + let freed = client.evict_registrations_for(base); + rail.stats.registered_bytes.fetch_sub(freed as u64, Ordering::Relaxed); let _ = ack.send(()); } Command::Stop => break, @@ -1220,6 +1248,13 @@ impl MultiRailClient { connections_quiesced: s.connections_quiesced.load(Ordering::Relaxed), inflight_requests: s.inflight_requests.load(Ordering::Relaxed), inflight_bytes: s.inflight_bytes.load(Ordering::Relaxed), + latency_avg_us: s + .latency_us_sum + .load(Ordering::Relaxed) + .checked_div(s.requests_ok.load(Ordering::Relaxed)) + .unwrap_or(0), + latency_max_us: s.latency_us_max.load(Ordering::Relaxed), + registered_bytes: s.registered_bytes.load(Ordering::Relaxed), } }) .collect() diff --git a/kv-service/client-rs/src/rdma.rs b/kv-service/client-rs/src/rdma.rs index 56a3fdd..ed8df47 100644 --- a/kv-service/client-rs/src/rdma.rs +++ b/kv-service/client-rs/src/rdma.rs @@ -687,13 +687,18 @@ impl RdmaClient { } /// Drop every cached registration whose base pointer is `base`, - /// deregistering the memory regions. Returns how many entries were - /// evicted. Multi-rail callers use this to quiesce registrations of a - /// caller buffer synchronously before returning it. + /// deregistering the memory regions. Returns how many bytes were + /// deregistered. Multi-rail callers use this to quiesce registrations of + /// a caller buffer synchronously before returning it. pub fn evict_registrations_for(&mut self, base: usize) -> usize { - let before = self.mr_cache.len(); + let before = self.mr_cache.iter().map(|((ptr, len), _)| if *ptr == base { *len } else { 0 }).sum(); self.mr_cache.retain(|((ptr, _), _)| *ptr != base); - before - self.mr_cache.len() + before + } + + /// Total bytes currently pinned by cached registrations. + pub fn registered_cache_bytes(&self) -> usize { + self.mr_cache.iter().map(|((_, len), _)| *len).sum() } /// Stripe-subset GET with a scatter destination list (wire tag 15): the diff --git a/kv-service/rdma-ffi/src/multirail.rs b/kv-service/rdma-ffi/src/multirail.rs index 414fbe3..586f044 100644 --- a/kv-service/rdma-ffi/src/multirail.rs +++ b/kv-service/rdma-ffi/src/multirail.rs @@ -104,6 +104,9 @@ pub struct CsMrRailStats { pub connections_quiesced: u64, pub inflight_requests: u64, pub inflight_bytes: u64, + pub latency_avg_us: u64, + pub latency_max_us: u64, + pub registered_bytes: u64, pub device: [c_char; 64], pub topology: [c_char; 160], } @@ -303,6 +306,9 @@ pub unsafe extern "C" fn cs_mr_rail_stats( connections_quiesced: snapshot.connections_quiesced, inflight_requests: snapshot.inflight_requests, inflight_bytes: snapshot.inflight_bytes, + latency_avg_us: snapshot.latency_avg_us, + latency_max_us: snapshot.latency_max_us, + registered_bytes: snapshot.registered_bytes, device: [0; 64], topology: [0; 160], }; diff --git a/src/contextstore/storage/multirail_client.py b/src/contextstore/storage/multirail_client.py index 8d5a2cf..b56b86f 100644 --- a/src/contextstore/storage/multirail_client.py +++ b/src/contextstore/storage/multirail_client.py @@ -76,6 +76,9 @@ class _CsMrRailStats(ctypes.Structure): ("connections_quiesced", ctypes.c_uint64), ("inflight_requests", ctypes.c_uint64), ("inflight_bytes", ctypes.c_uint64), + ("latency_avg_us", ctypes.c_uint64), + ("latency_max_us", ctypes.c_uint64), + ("registered_bytes", ctypes.c_uint64), ("device", ctypes.c_char * 64), ("topology", ctypes.c_char * 160), ] @@ -96,6 +99,9 @@ class RailStats: connections_quiesced: int inflight_requests: int inflight_bytes: int + latency_avg_us: int + latency_max_us: int + registered_bytes: int def __str__(self) -> str: return ( @@ -104,7 +110,9 @@ def __str__(self) -> str: f"ok={self.requests_ok} err={self.requests_err} tmo={self.timeouts} " f"read={self.bytes_read >> 20}MiB conns=+{self.connections_created}" f"/~{self.connections_quiesced} " - f"inflight=(req:{self.inflight_requests},B:{self.inflight_bytes})" + f"inflight=(req:{self.inflight_requests},B:{self.inflight_bytes}) " + f"lat=(avg:{self.latency_avg_us}us,max:{self.latency_max_us}us) " + f"reg={self.registered_bytes >> 20}MiB" ) @@ -253,6 +261,9 @@ def rail_stats(self) -> list[RailStats]: connections_quiesced=entry.connections_quiesced, inflight_requests=entry.inflight_requests, inflight_bytes=entry.inflight_bytes, + latency_avg_us=entry.latency_avg_us, + latency_max_us=entry.latency_max_us, + registered_bytes=entry.registered_bytes, ) ) return stats From ceffb4819429eebe1a4ba86befb79f8598eb848e Mon Sep 17 00:00:00 2001 From: wzu Date: Sat, 26 Sep 2026 13:53:20 +0800 Subject: [PATCH 5/7] =?UTF-8?q?multirail:=20PR=20hardening=20=E2=80=94=20c?= =?UTF-8?q?ounter-leak=20fixes,=20rustfmt=20pass?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Self-review findings fixed before opening the upstream PR: - aborted reads now drain outstanding task replies (bounded by the read deadline) before quiescing, releasing inflight_requests/inflight_bytes — previously an aborted wave could leak counters and eventually starve the backpressure wait; - a connection worker subtracts its still-cached registration bytes from the rail registered_bytes counter on teardown, keeping the metric honest after failures; - cargo fmt applied to the touched crates (reverted on untouched upstream files to keep the diff noise-free). --- .../client-rs/src/bin/multirail_bench.rs | 45 ++-- kv-service/client-rs/src/bin/rdma_bench.rs | 7 +- kv-service/client-rs/src/lib.rs | 6 +- kv-service/client-rs/src/multirail.rs | 215 +++++++++++++----- kv-service/client-rs/src/rdma.rs | 6 +- kv-service/client-rs/tests/multirail_e2e.rs | 73 +++--- kv-service/rdma-ffi/src/multirail.rs | 27 ++- 7 files changed, 257 insertions(+), 122 deletions(-) diff --git a/kv-service/client-rs/src/bin/multirail_bench.rs b/kv-service/client-rs/src/bin/multirail_bench.rs index 37509a1..500cbdc 100644 --- a/kv-service/client-rs/src/bin/multirail_bench.rs +++ b/kv-service/client-rs/src/bin/multirail_bench.rs @@ -188,16 +188,17 @@ fn format_coordinator(url: &str) -> String { } fn lookup(args: &Args, runtime: &tokio::runtime::Runtime) -> Result { - runtime.block_on(async { - let mut client = KvClient::connect(format_coordinator(&args.coordinator)) - .await - .map_err(|error| anyhow!(error.to_string()))?; - client - .lookup_object(&args.namespace, &args.object_key) - .await - .map_err(|error| anyhow!(error.to_string())) - })? - .ok_or_else(|| anyhow!("object not found: {}/{}", args.namespace, args.object_key)) + runtime + .block_on(async { + let mut client = KvClient::connect(format_coordinator(&args.coordinator)) + .await + .map_err(|error| anyhow!(error.to_string()))?; + client + .lookup_object(&args.namespace, &args.object_key) + .await + .map_err(|error| anyhow!(error.to_string())) + })? + .ok_or_else(|| anyhow!("object not found: {}/{}", args.namespace, args.object_key)) } fn run_client( @@ -213,7 +214,9 @@ fn run_client( task_max_stripes: args.task_max_stripes, ..RailLimits::default() }; - let client = MultiRailClient::new(rails)?.with_limits(limits).with_policy(policy); + let client = MultiRailClient::new(rails)? + .with_limits(limits) + .with_policy(policy); let mut buffer = AlignedBuffer::new(args.buf_mb * 1024 * 1024)?; let object_size = usize::try_from(lookup.descriptor.size)?; @@ -241,7 +244,11 @@ fn run_client( unsafe { client.read_object_into_raw( &lookup.descriptor, - &lookup.placement.as_ref().map(|p| p.chunks.clone()).unwrap_or_default(), + &lookup + .placement + .as_ref() + .map(|p| p.chunks.clone()) + .unwrap_or_default(), buffer.ptr, buffer.len, true, @@ -259,7 +266,10 @@ fn run_client( )); } if args.verify { - verify_pattern(&buffer.as_mut()[..object_size], &format!("{label}#{iteration}"))?; + verify_pattern( + &buffer.as_mut()[..object_size], + &format!("{label}#{iteration}"), + )?; } } @@ -267,7 +277,10 @@ fn run_client( println!("[{label}] {snapshot}"); } let total: f64 = latencies.iter().map(|d| d.as_secs_f64()).sum(); - let best = latencies.iter().map(|d| d.as_secs_f64()).fold(f64::INFINITY, f64::min); + let best = latencies + .iter() + .map(|d| d.as_secs_f64()) + .fold(f64::INFINITY, f64::min); let gbps = object_size as f64 / 1024f64.powi(3) / (total / latencies.len() as f64); println!( "[{label}] avg={:.3}s best={:.3}s avg_bw={:.3} GiB/s over {} iterations", @@ -308,7 +321,9 @@ fn main() -> Result<()> { } let mut lookup = lookup(&args, &runtime)?; if lookup.placement.is_none() { - return Err(anyhow!("lookup returned no placement (is the object striped?)")); + return Err(anyhow!( + "lookup returned no placement (is the object striped?)" + )); } let original_lookup = lookup.clone(); // Reachability indirection: rewrite placement endpoints through the map diff --git a/kv-service/client-rs/src/bin/rdma_bench.rs b/kv-service/client-rs/src/bin/rdma_bench.rs index 85ac6ce..ccba6fd 100644 --- a/kv-service/client-rs/src/bin/rdma_bench.rs +++ b/kv-service/client-rs/src/bin/rdma_bench.rs @@ -161,7 +161,9 @@ fn main() -> Result<()> { /// WRITE + commit + striped O_DIRECT pwrite) instead of dedup-skipping. fn run_put(args: &Args) -> Result<()> { if args.coordinator.is_some() { - return Err(anyhow!("--mode put does not support --coordinator (single endpoint only)")); + return Err(anyhow!( + "--mode put does not support --coordinator (single endpoint only)" + )); } let namespace = args .namespace @@ -697,8 +699,7 @@ fn run_multi_endpoint(args: &Args, coordinator: &str) -> Result<()> { let seg_len = (buf_size / sge_segments.max(1)).max(1) as u64; let mut segments = Vec::with_capacity(sge_segments); let mut off = 0u64; - let (base, rkey, total) = - (view.addr(), view.rkey(), buf_size as u64); + let (base, rkey, total) = (view.addr(), view.rkey(), buf_size as u64); while off < total { let n = seg_len.min(total - off); segments.push((base + off, rkey, n)); diff --git a/kv-service/client-rs/src/lib.rs b/kv-service/client-rs/src/lib.rs index ef0d675..98c2748 100644 --- a/kv-service/client-rs/src/lib.rs +++ b/kv-service/client-rs/src/lib.rs @@ -634,9 +634,9 @@ impl KvClient { let offset = usize::try_from(chunk.offset).map_err(|_| { tonic::Status::internal(format!("negative chunk offset {}", chunk.offset)) })?; - let end = offset.checked_add(chunk.data.len()).ok_or_else(|| { - tonic::Status::internal("chunk offset + length overflows usize") - })?; + let end = offset + .checked_add(chunk.data.len()) + .ok_or_else(|| tonic::Status::internal("chunk offset + length overflows usize"))?; if end > dst.len() { return Err(tonic::Status::internal(format!( "chunk [{offset}, {end}) exceeds destination buffer of {} bytes", diff --git a/kv-service/client-rs/src/multirail.rs b/kv-service/client-rs/src/multirail.rs index 3422a6b..1fe5204 100644 --- a/kv-service/client-rs/src/multirail.rs +++ b/kv-service/client-rs/src/multirail.rs @@ -152,7 +152,10 @@ impl RailConfig { if self.endpoints.is_empty() { return true; } - let host = endpoint.rsplit_once(':').map(|(h, _)| h).unwrap_or(endpoint); + let host = endpoint + .rsplit_once(':') + .map(|(h, _)| h) + .unwrap_or(endpoint); self.endpoints .iter() .any(|allowed| allowed == endpoint || *allowed == host) @@ -270,12 +273,20 @@ impl RailSelectPolicy { pub enum MultiRailError { NoRails, NoHealthyRails, - BufferTooSmall { need: u64, have: usize }, + BufferTooSmall { + need: u64, + have: usize, + }, InvalidPlacement(String), MissingStripes(Vec), DuplicateStripes(Vec), - OutOfBoundsStripe { stripe: u32, stripe_count: u32 }, - StaleDescriptor { endpoint: String }, + OutOfBoundsStripe { + stripe: u32, + stripe_count: u32, + }, + StaleDescriptor { + endpoint: String, + }, TaskFailed { rail: String, endpoint: String, @@ -303,7 +314,10 @@ pub enum MultiRailError { expected: String, actual: String, }, - WorkerPanic { rail: String, endpoint: String }, + WorkerPanic { + rail: String, + endpoint: String, + }, Cancelled, } @@ -460,7 +474,10 @@ fn topology_from(base: &Path, device: &str) -> RailTopology { .map(|name| name.to_string_lossy().into_owned()) }) .unwrap_or_default(); - RailTopology { numa_node: numa, pci_slot } + RailTopology { + numa_node: numa, + pci_slot, + } } /// Read a rail's topology from the real sysfs tree. @@ -648,7 +665,10 @@ fn validate_placement( for chunk in chunks { let stripe = chunk.stripe_index; if stripe >= stripe_count { - return Err(MultiRailError::OutOfBoundsStripe { stripe, stripe_count }); + return Err(MultiRailError::OutOfBoundsStripe { + stripe, + stripe_count, + }); } if chunk.rdma_endpoint.is_empty() { return Err(MultiRailError::InvalidPlacement(format!( @@ -657,7 +677,10 @@ fn validate_placement( } let expected_len = if descriptor.is_striped { stripe_length(stripe, stripe_count, chunk_size, object_size).ok_or( - MultiRailError::OutOfBoundsStripe { stripe, stripe_count }, + MultiRailError::OutOfBoundsStripe { + stripe, + stripe_count, + }, )? } else { object_size @@ -769,9 +792,7 @@ fn least_loaded(candidates: &[usize], rails: &[Arc], assigned: &[u64]) -> // assigned ≤ 2^40 and weight ≤ 2^32, so the products fit u128. let lhs = (assigned[candidate] as u128) * (rails[best].config.weight as u128); let rhs = (assigned[best] as u128) * (rails[candidate].config.weight as u128); - if lhs < rhs - || (lhs == rhs && rails[candidate].config.weight > rails[best].config.weight) - { + if lhs < rhs || (lhs == rhs && rails[candidate].config.weight > rails[best].config.weight) { best = candidate; } } @@ -815,7 +836,8 @@ fn pick_rail( .iter() .copied() .filter(|&index| { - rails[index].gid_v4 + rails[index] + .gid_v4 .is_some_and(|rail_ip| rail_ip.octets()[..3] == ip.octets()[..3]) }) .collect(); @@ -1020,10 +1042,7 @@ struct ConnEntry { } /// Expected bytes/chunks for one stripe-subset request against `descriptor`. -fn expected_task_outcome( - descriptor: &pb::ObjectDescriptor, - stripes: &[u32], -) -> (u64, u32) { +fn expected_task_outcome(descriptor: &pb::ObjectDescriptor, stripes: &[u32]) -> (u64, u32) { if stripes.is_empty() { // Non-striped (or whole-object) GET: one implicit chunk. return (descriptor.size, 1); @@ -1076,7 +1095,9 @@ fn conn_worker( ) { let mut client = match connect_rail_client(&rail, &endpoint, &limits) { Ok(client) => { - rail.stats.connections_created.fetch_add(1, Ordering::Relaxed); + rail.stats + .connections_created + .fetch_add(1, Ordering::Relaxed); let _ = setup_tx.send(Ok(())); client } @@ -1119,10 +1140,9 @@ fn conn_worker( }; let registered_now = client.registered_cache_bytes(); if registered_now > before { - rail.stats.registered_bytes.fetch_add( - (registered_now - before) as u64, - Ordering::Relaxed, - ); + rail.stats + .registered_bytes + .fetch_add((registered_now - before) as u64, Ordering::Relaxed); } client .get_descriptor_stripes_into_view_detailed(&descriptor, &stripes, view, 0) @@ -1131,10 +1151,9 @@ fn conn_worker( rail.stats .latency_us_sum .fetch_add(started.elapsed().as_micros() as u64, Ordering::Relaxed); - rail.stats.latency_us_max.fetch_max( - started.elapsed().as_micros() as u64, - Ordering::Relaxed, - ); + rail.stats + .latency_us_max + .fetch_max(started.elapsed().as_micros() as u64, Ordering::Relaxed); let _ = reply.send(TaskReply { rail_index: rail.index, endpoint: Arc::clone(&endpoint), @@ -1145,12 +1164,20 @@ fn conn_worker( } Command::EvictRegistration { base, ack } => { let freed = client.evict_registrations_for(base); - rail.stats.registered_bytes.fetch_sub(freed as u64, Ordering::Relaxed); + rail.stats + .registered_bytes + .fetch_sub(freed as u64, Ordering::Relaxed); let _ = ack.send(()); } Command::Stop => break, } } + + // Teardown bookkeeping: whatever this connection still had registered + // goes away with the dropped client — keep the rail counter honest. + rail.stats + .registered_bytes + .fetch_sub(client.registered_cache_bytes() as u64, Ordering::Relaxed); } // --------------------------------------------------------------------------- @@ -1181,9 +1208,8 @@ impl MultiRailClient { .enumerate() .map(|(index, config)| { let topology = read_topology(&config.device); - let gid_v4 = - rdma::query_gid_raw(&config.device, config.port, config.gid_index) - .and_then(gid_to_v4); + let gid_v4 = rdma::query_gid_raw(&config.device, config.port, config.gid_index) + .and_then(gid_to_v4); Arc::new(Rail { index, topology, @@ -1325,7 +1351,14 @@ impl MultiRailClient { have: len, }); } - self.read_impl(descriptor, chunks, ptr as usize, len, sticky_registration, None) + self.read_impl( + descriptor, + chunks, + ptr as usize, + len, + sticky_registration, + None, + ) } fn read_impl( @@ -1366,10 +1399,7 @@ impl MultiRailClient { // Global deadline: every wave gets a full io_timeout, plus one extra // for connect phases. let deadline = Instant::now() - + self - .limits - .io_timeout - .mul_f64(waves.len().max(1) as f64) + + self.limits.io_timeout.mul_f64(waves.len().max(1) as f64) + self.limits.io_timeout; let (reply_tx, reply_rx) = mpsc::channel::(); @@ -1379,6 +1409,8 @@ impl MultiRailClient { let mut failed_rail: Option = None; let mut verified_bytes = 0u64; + let mut dispatched_total = 0usize; + let mut collected_total = 0usize; 'waves: for wave in waves { if let Some(cancel) = cancel { cancel.check()?; @@ -1403,7 +1435,9 @@ impl MultiRailClient { Ok(tx) => { participated.insert((rail.index, Arc::clone(&task.endpoint))); rail.stats.inflight_requests.fetch_add(1, Ordering::Relaxed); - rail.stats.inflight_bytes.fetch_add(task.bytes, Ordering::Relaxed); + rail.stats + .inflight_bytes + .fetch_add(task.bytes, Ordering::Relaxed); let sent = tx.send(Command::Read { descriptor: Arc::clone(&descriptor), stripes: Arc::new(task.stripes.clone()), @@ -1413,7 +1447,9 @@ impl MultiRailClient { }); if sent.is_err() { rail.stats.inflight_requests.fetch_sub(1, Ordering::Relaxed); - rail.stats.inflight_bytes.fetch_sub(task.bytes, Ordering::Relaxed); + rail.stats + .inflight_bytes + .fetch_sub(task.bytes, Ordering::Relaxed); failure = Some(MultiRailError::TaskFailed { rail: rail.config.device.clone(), endpoint: task.endpoint.to_string(), @@ -1423,6 +1459,7 @@ impl MultiRailClient { break 'waves; } dispatched += 1; + dispatched_total += 1; } Err(error) => { failed_rail = Some(rail.index); @@ -1456,7 +1493,10 @@ impl MultiRailClient { // Bound each wait so cancellation is observed promptly. let slice = remaining.min(Duration::from_millis(100)); match reply_rx.recv_timeout(slice) { - Ok(reply) => replies.push(reply), + Ok(reply) => { + replies.push(reply); + collected_total += 1; + } Err(mpsc::RecvTimeoutError::Timeout) => { if deadline <= Instant::now() { failure = Some(MultiRailError::Timeout { @@ -1549,7 +1589,31 @@ impl MultiRailClient { } } - // ---- failure path: quiesce participating connections ---- + // ---- failure path: drain outstanding replies, then quiesce ---- + // Workers send exactly one reply per Read, and the quiesce Stop is + // queued behind any in-flight command — so every outstanding reply + // arrives once the worker finishes. Drain them (bounded by the read + // deadline) to release the in-flight counters; otherwise aborted + // reads would leak inflight_requests/inflight_bytes and eventually + // starve await_headroom. + while collected_total < dispatched_total { + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + break; + } + match reply_rx.recv_timeout(remaining.min(Duration::from_millis(500))) { + Ok(reply) => { + collected_total += 1; + let rail = &self.rails[reply.rail_index]; + rail.stats.inflight_requests.fetch_sub(1, Ordering::Relaxed); + rail.stats + .inflight_bytes + .fetch_sub(reply.expected_bytes, Ordering::Relaxed); + } + Err(_) => break, + } + } + if let Some(error) = failure { if let Some(index) = failed_rail { self.rails[index].mark_unhealthy(self.limits.rail_cooldown); @@ -1582,8 +1646,9 @@ impl MultiRailClient { } // SAFETY: `base..base+len` is the caller buffer, valid for // this whole call; the slice below stays inside it. - let view: &[u8] = - unsafe { std::slice::from_raw_parts((base + start) as *const u8, info.length as usize) }; + let view: &[u8] = unsafe { + std::slice::from_raw_parts((base + start) as *const u8, info.length as usize) + }; let actual = format!("{:016x}", twox_hash::xxh3::hash64(view)); if !expected.eq_ignore_ascii_case(&actual) { return Err(MultiRailError::ChecksumMismatch { @@ -1666,7 +1731,12 @@ impl MultiRailClient { endpoint: Arc, deadline: Instant, ) -> Result, MultiRailError> { - if let Some(entry) = self.conns.lock().unwrap().get(&(rail.index, Arc::clone(&endpoint))) { + if let Some(entry) = self + .conns + .lock() + .unwrap() + .get(&(rail.index, Arc::clone(&endpoint))) + { return Ok(entry.tx.clone()); } @@ -1677,9 +1747,7 @@ impl MultiRailClient { let limits = self.limits.clone(); let handle = std::thread::Builder::new() .name(format!("mrail-{}-{}", rail.config.device, endpoint)) - .spawn(move || { - conn_worker(worker_rail, worker_endpoint, limits, command_rx, setup_tx) - }) + .spawn(move || conn_worker(worker_rail, worker_endpoint, limits, command_rx, setup_tx)) .map_err(|error| MultiRailError::TaskFailed { rail: rail.config.device.clone(), endpoint: endpoint.to_string(), @@ -1735,7 +1803,9 @@ impl MultiRailClient { fn quiesce(&self, keys: Vec<(usize, Arc)>) { let entries: Vec = { let mut conns = self.conns.lock().unwrap(); - keys.into_iter().filter_map(|key| conns.remove(&key)).collect() + keys.into_iter() + .filter_map(|key| conns.remove(&key)) + .collect() }; for entry in &entries { let _ = entry.tx.send(Command::Stop); @@ -1925,7 +1995,13 @@ mod tests { let chunks: Vec<_> = (0..8).map(|i| chunk(i, "10.0.0.1", 8, "")).collect(); let placement = validate_placement(&desc, &chunks).unwrap(); let rails = test_rails(2); - let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded, &RailLimits::default()).unwrap(); + let plan = build_plan( + &placement, + &rails, + RailSelectPolicy::LeastLoaded, + &RailLimits::default(), + ) + .unwrap(); assert_eq!(plan.task_count(), 2); let total: u64 = plan.tasks.iter().map(|t| t.bytes).sum(); assert_eq!(total, 64); @@ -1959,7 +2035,13 @@ mod tests { unhealthy_until: Mutex::new(None), gid_v4: None, }); - let plan = build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded, &RailLimits::default()).unwrap(); + let plan = build_plan( + &placement, + &rails, + RailSelectPolicy::LeastLoaded, + &RailLimits::default(), + ) + .unwrap(); let rail0: u64 = plan .tasks .iter() @@ -1991,12 +2073,21 @@ mod tests { unhealthy_until: Mutex::new(None), gid_v4: None, }); - let plan = - build_plan(&placement, &rails, RailSelectPolicy::LeastLoaded, &RailLimits::default()).expect("plan builds"); + let plan = build_plan( + &placement, + &rails, + RailSelectPolicy::LeastLoaded, + &RailLimits::default(), + ) + .expect("plan builds"); // dev0 pinned to 10.0.0.2; 10.0.0.1 falls to dev1. assert_eq!(plan.tasks.len(), 2); for task in &plan.tasks { - let expected_rail = if task.endpoint.contains("10.0.0.2") { 0 } else { 1 }; + let expected_rail = if task.endpoint.contains("10.0.0.2") { + 0 + } else { + 1 + }; assert_eq!(task.rail_index, expected_rail); } // No rail allowed for 10.0.0.1 → typed planning error. @@ -2010,7 +2101,13 @@ mod tests { unhealthy_until: Mutex::new(None), gid_v4: None, }); - assert!(build_plan(&placement, &strict, RailSelectPolicy::LeastLoaded, &RailLimits::default()).is_err()); + assert!(build_plan( + &placement, + &strict, + RailSelectPolicy::LeastLoaded, + &RailLimits::default() + ) + .is_err()); } #[test] @@ -2019,7 +2116,13 @@ mod tests { let chunks: Vec<_> = (0..8).map(|i| chunk(i, "10.0.0.1", 8, "")).collect(); let placement = validate_placement(&desc, &chunks).unwrap(); let rails = test_rails(2); - let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity, &RailLimits::default()).unwrap(); + let plan = build_plan( + &placement, + &rails, + RailSelectPolicy::EndpointAffinity, + &RailLimits::default(), + ) + .unwrap(); assert_eq!(plan.task_count(), 1); assert_eq!(plan.tasks[0].bytes, 64); assert_eq!(plan.tasks[0].stripes.len(), 8); @@ -2040,7 +2143,13 @@ mod tests { unhealthy_until: Mutex::new(None), gid_v4: Some("10.0.0.5".parse().unwrap()), }); - let plan = build_plan(&placement, &rails, RailSelectPolicy::EndpointAffinity, &RailLimits::default()).unwrap(); + let plan = build_plan( + &placement, + &rails, + RailSelectPolicy::EndpointAffinity, + &RailLimits::default(), + ) + .unwrap(); assert_eq!(plan.tasks.len(), 1); assert_eq!(plan.tasks[0].rail_index, 1); } diff --git a/kv-service/client-rs/src/rdma.rs b/kv-service/client-rs/src/rdma.rs index ed8df47..8207ee1 100644 --- a/kv-service/client-rs/src/rdma.rs +++ b/kv-service/client-rs/src/rdma.rs @@ -691,7 +691,11 @@ impl RdmaClient { /// deregistered. Multi-rail callers use this to quiesce registrations of /// a caller buffer synchronously before returning it. pub fn evict_registrations_for(&mut self, base: usize) -> usize { - let before = self.mr_cache.iter().map(|((ptr, len), _)| if *ptr == base { *len } else { 0 }).sum(); + let before = self + .mr_cache + .iter() + .map(|((ptr, len), _)| if *ptr == base { *len } else { 0 }) + .sum(); self.mr_cache.retain(|((ptr, _), _)| *ptr != base); before } diff --git a/kv-service/client-rs/tests/multirail_e2e.rs b/kv-service/client-rs/tests/multirail_e2e.rs index a1985b9..bd4c888 100644 --- a/kv-service/client-rs/tests/multirail_e2e.rs +++ b/kv-service/client-rs/tests/multirail_e2e.rs @@ -11,9 +11,7 @@ #![cfg(feature = "rdma")] -use contextstore_client_rs::multirail::{ - MultiRailClient, MultiRailError, RailConfig, RailLimits, -}; +use contextstore_client_rs::multirail::{MultiRailClient, MultiRailError, RailConfig, RailLimits}; use contextstore_client_rs::{KvClient, ObjectLookup}; use prost::bytes::Bytes; use std::time::{Duration, SystemTime, UNIX_EPOCH}; @@ -136,27 +134,26 @@ fn seed(namespace: &str, key: &str, size_mb: usize) -> Fixture { let mut payload = vec![0u8; size]; fill_pattern(&mut payload); let coordinator = env_or("CS_MR_COORDINATOR", "http://127.0.0.1:50051"); - let lookup = runtime - .block_on(async { - let mut client = KvClient::connect(coordinator) - .await - .expect("connect gRPC coordinator"); - let big = Bytes::from(payload); - let chunk = 4 * 1024 * 1024; - let mut segments = Vec::new(); - for offset in (0..size).step_by(chunk) { - segments.push(big.slice(offset..(offset + chunk).min(size))); - } - client - .put_stream_chunks(namespace, key, segments) - .await - .expect("seed object over gRPC"); - client - .lookup_object(namespace, key) - .await - .expect("lookup seeded object") - .expect("seeded object present") - }); + let lookup = runtime.block_on(async { + let mut client = KvClient::connect(coordinator) + .await + .expect("connect gRPC coordinator"); + let big = Bytes::from(payload); + let chunk = 4 * 1024 * 1024; + let mut segments = Vec::new(); + for offset in (0..size).step_by(chunk) { + segments.push(big.slice(offset..(offset + chunk).min(size))); + } + client + .put_stream_chunks(namespace, key, segments) + .await + .expect("seed object over gRPC"); + client + .lookup_object(namespace, key) + .await + .expect("lookup seeded object") + .expect("seeded object present") + }); Fixture { runtime, namespace: namespace.to_string(), @@ -286,21 +283,19 @@ fn stale_descriptor_is_reported_as_relookup() { let mut payload = vec![0u8; fixture.size]; fill_pattern(&mut payload); payload[0] ^= 0xFF; // different content → different etag/generation - fixture - .runtime - .block_on(async { - let mut client = KvClient::connect(coordinator) - .await - .expect("connect coordinator"); - client - .delete(&fixture.namespace, &fixture.key) - .await - .expect("delete old version"); - client - .put(&fixture.namespace, &fixture.key, payload) - .await - .expect("rewrite object"); - }); + fixture.runtime.block_on(async { + let mut client = KvClient::connect(coordinator) + .await + .expect("connect coordinator"); + client + .delete(&fixture.namespace, &fixture.key) + .await + .expect("delete old version"); + client + .put(&fixture.namespace, &fixture.key, payload) + .await + .expect("rewrite object"); + }); let client = MultiRailClient::new(rails_from_env()) .expect("create client") diff --git a/kv-service/rdma-ffi/src/multirail.rs b/kv-service/rdma-ffi/src/multirail.rs index 586f044..a0cf180 100644 --- a/kv-service/rdma-ffi/src/multirail.rs +++ b/kv-service/rdma-ffi/src/multirail.rs @@ -128,9 +128,9 @@ unsafe fn cstr<'a>(ptr: *const c_char) -> Result<&'a str, MultiRailError> { if ptr.is_null() { return Err(MultiRailError::InvalidPlacement("null string".into())); } - CStr::from_ptr(ptr).to_str().map_err(|_| { - MultiRailError::InvalidPlacement("non-utf8 string in descriptor".into()) - }) + CStr::from_ptr(ptr) + .to_str() + .map_err(|_| MultiRailError::InvalidPlacement("non-utf8 string in descriptor".into())) } #[no_mangle] @@ -141,7 +141,12 @@ pub unsafe extern "C" fn cs_mr_new( ) -> *mut c_void { let mut rails = Vec::with_capacity(rail_count as usize); for index in 0..rail_count as usize { - let spec = match rail_specs.add(index).read().as_ref().and_then(|p| CStr::from_ptr(p).to_str().ok()) { + let spec = match rail_specs + .add(index) + .read() + .as_ref() + .and_then(|p| CStr::from_ptr(p).to_str().ok()) + { Some(spec) => spec, None => return std::ptr::null_mut(), }; @@ -156,7 +161,9 @@ pub unsafe extern "C" fn cs_mr_new( }; match MultiRailClient::new(rails) { Ok(client) => Box::into_raw(Box::new(MrReader { - client: client.with_limits(limits).with_policy(RailSelectPolicy::LeastLoaded), + client: client + .with_limits(limits) + .with_policy(RailSelectPolicy::LeastLoaded), })) as *mut c_void, Err(_) => std::ptr::null_mut(), } @@ -252,9 +259,13 @@ pub unsafe extern "C" fn cs_mr_read( } else { // SAFETY: the caller guarantees buffer..buffer+len stays valid and // unmoved for the duration of the call (it is synchronous). - reader - .client - .read_object_into_raw(&pb_descriptor, &pb_chunks, buffer, buffer_len as usize, false) + reader.client.read_object_into_raw( + &pb_descriptor, + &pb_chunks, + buffer, + buffer_len as usize, + false, + ) }; match result { Ok(bytes) => bytes as i64, From c155086129cfab0c23db5b294f88d8b4801c93e1 Mon Sep 17 00:00:00 2001 From: 123123213weqw <123123213weqw@users.noreply.github.com> Date: Sat, 26 Sep 2026 15:07:56 +0800 Subject: [PATCH 6/7] multirail: unit test for cooperative CancelToken semantics --- kv-service/client-rs/src/multirail.rs | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/kv-service/client-rs/src/multirail.rs b/kv-service/client-rs/src/multirail.rs index 1fe5204..314a5b9 100644 --- a/kv-service/client-rs/src/multirail.rs +++ b/kv-service/client-rs/src/multirail.rs @@ -2205,6 +2205,18 @@ mod tests { assert_eq!(endpoint_v4("not-a-host:50053"), None); } + #[test] + fn cancel_token_is_shared_and_maps_to_cancelled() { + let token = CancelToken::new(); + let observer = token.clone(); + assert!(!token.is_cancelled()); + // 取消在克隆之间共享同一标志; wave 边界/背压等待/应答收集点 + // 经 check() 协作退出, 映射为类型化 Cancelled 错误。 + observer.cancel(); + assert!(token.is_cancelled()); + assert!(matches!(token.check(), Err(MultiRailError::Cancelled))); + } + #[test] fn checksum_verification_matches_server_encoding() { let mut buffer = vec![0u8; 16]; From f122e83c04c543a3123b3c35abc2730ec6879344 Mon Sep 17 00:00:00 2001 From: 123123213weqw <123123213weqw@users.noreply.github.com> Date: Fri, 9 Oct 2026 16:53:18 +0800 Subject: [PATCH 7/7] multirail: evict non-sticky MRs on verification error paths Review feedback: the ByteCountMismatch and ChecksumMismatch early returns skipped the synchronous registration eviction that the success path (and the quiesce-backed failure path) perform, so a cached MR could outlive the caller's buffer and pin its freed pages until connection teardown. Extract evict_participated_registrations() and run it on every post-dispatch error return. --- kv-service/client-rs/src/multirail.rs | 64 +++++++++++++++++---------- 1 file changed, 41 insertions(+), 23 deletions(-) diff --git a/kv-service/client-rs/src/multirail.rs b/kv-service/client-rs/src/multirail.rs index 314a5b9..366aa22 100644 --- a/kv-service/client-rs/src/multirail.rs +++ b/kv-service/client-rs/src/multirail.rs @@ -1625,6 +1625,9 @@ impl MultiRailClient { // Full-coverage invariant: the per-task byte checks must add up to // the whole object, otherwise stripes went missing. if verified_bytes != placement.expected_bytes { + if !sticky { + self.evict_participated_registrations(&participated, base); + } return Err(MultiRailError::ByteCountMismatch { rail: "aggregate".into(), endpoint: "aggregate".into(), @@ -1651,6 +1654,9 @@ impl MultiRailClient { }; let actual = format!("{:016x}", twox_hash::xxh3::hash64(view)); if !expected.eq_ignore_ascii_case(&actual) { + if !sticky { + self.evict_participated_registrations(&participated, base); + } return Err(MultiRailError::ChecksumMismatch { stripe: *stripe, expected: expected.clone(), @@ -1662,29 +1668,7 @@ impl MultiRailClient { // ---- success path: synchronous MR eviction for non-sticky buffers ---- if !sticky { - for key in &participated { - let ack_rx = { - let conns = self.conns.lock().unwrap(); - match conns.get(key) { - Some(entry) => { - let (ack_tx, ack_rx) = mpsc::channel(); - if entry - .tx - .send(Command::EvictRegistration { base, ack: ack_tx }) - .is_ok() - { - Some(ack_rx) - } else { - None - } - } - None => None, - } - }; - if let Some(ack_rx) = ack_rx { - let _ = ack_rx.recv_timeout(self.limits.io_timeout); - } - } + self.evict_participated_registrations(&participated, base); } Ok(placement.expected_bytes as usize) @@ -1800,6 +1784,40 @@ impl MultiRailClient { /// completion implies: worker loop exited → `RdmaClient` dropped → BYE /// sent, QP destroyed, MRs deregistered. This is the quiesce barrier that /// makes returning a failed read safe. + /// Evict the caller-buffer registrations from all participating + /// connections (non-sticky reads). Runs on the success path AND on + /// post-dispatch error returns — a cached MR must never outlive the + /// caller's buffer and pin its freed pages. + fn evict_participated_registrations( + &self, + participated: &std::collections::HashSet<(usize, Arc)>, + base: usize, + ) { + for key in participated { + let ack_rx = { + let conns = self.conns.lock().unwrap(); + match conns.get(key) { + Some(entry) => { + let (ack_tx, ack_rx) = mpsc::channel(); + if entry + .tx + .send(Command::EvictRegistration { base, ack: ack_tx }) + .is_ok() + { + Some(ack_rx) + } else { + None + } + } + None => None, + } + }; + if let Some(ack_rx) = ack_rx { + let _ = ack_rx.recv_timeout(self.limits.io_timeout); + } + } + } + fn quiesce(&self, keys: Vec<(usize, Arc)>) { let entries: Vec = { let mut conns = self.conns.lock().unwrap();