Scale write throughput across data-shard groups while keeping a single unified read surface: - scatter_gather.rs: pooled fan-out across shard groups (replaces per-request client construction); cross-shard query results merged on one node - cluster/node.rs: cross-shard read routing — a read on any node gathers from every shard group's leader and unions results - cluster/forward.rs: fix h2 204 forward-relay bug (relay_forwarded skips body for 1xx/204/304 — synthesized JSON body on a 204 triggered HTTP/2 RST_STREAM on the real mTLS plane) - dto.rs: cross-shard query/result DTOs - k8s/cluster/: enable 3-group `shards:` topology (statefulset, service-peers, topology-configmap) - k8s/cluster-local-kind/: local-kind overlay to run the T5 gate without Ref-A - tidal-stress/k8s/stress-job-t5.yaml: 2-generator sharded throughput job - tests: cluster_cross_shard_reads.rs + multiproc support; ran real on kind - docs/profiling/m12p4-t5-sharded-throughput.md: T5 throughput findings
2494 lines
98 KiB
Rust
2494 lines
98 KiB
Rust
//! Scatter-gather query routing for entity-sharded cluster deployments.
|
||
//!
|
||
//! When entities are hash-partitioned across shards, a RETRIEVE or SEARCH
|
||
//! query must fan out to all shards, collect per-shard results, and merge
|
||
//! them into a single ranked result set. This module implements:
|
||
//!
|
||
//! - **Entity-sharded write routing**: `hash(entity_id) % num_shards`
|
||
//! - **Scatter-gather RETRIEVE**: fan out to all shards, K-way merge by score
|
||
//! - **Scatter-gather SEARCH**: fan out to all shards, merge by score
|
||
//! - **Deadline propagation**: configurable budget with per-hop overhead subtracted
|
||
//! - **Partial failure**: degraded results when some shards are unreachable
|
||
//!
|
||
//! # Size / cohesion note (M0–M10 review Maintainability-S)
|
||
//!
|
||
//! This file is large (~1.5k lines), but roughly half is its in-module test
|
||
//! suite and the production code is a single tightly-coupled concern: the
|
||
//! fan-out → gather → dedup → reconcile → diversity → truncate merge pipeline,
|
||
//! whose helpers share the [`Sourced`]/[`GatherState`]/[`MergeItem`] types and
|
||
//! must stay consistent. Splitting it would fragment that pipeline without
|
||
//! reducing real complexity, so the split is deferred (not forced mid-campaign).
|
||
//! Tracked for a dedicated follow-up.
|
||
|
||
use std::{
|
||
collections::{HashMap, HashSet},
|
||
sync::{Arc, OnceLock, mpsc},
|
||
time::{Duration, Instant},
|
||
};
|
||
|
||
use tidaldb::{
|
||
query::{
|
||
retrieve::{Results as RetrieveResults, Retrieve, RetrieveResult},
|
||
search::{Search, SearchResults},
|
||
},
|
||
replication::shard::{RegionId, ShardId, ShardRouter},
|
||
schema::EntityId,
|
||
testing::SimulatedCluster,
|
||
};
|
||
|
||
use crate::error::{Result, ServerError};
|
||
|
||
// ── Shard-executor seam ───────────────────────────────────────────────────────
|
||
|
||
/// The two cluster-fabric operations the scatter-gather merge pipeline depends on.
|
||
///
|
||
/// Abstracted so the SAME merge/dedup/diversity logic drives both the
|
||
/// single-process [`SimulatedCluster`] (in-process fetch + metadata) and the
|
||
/// multi-process region node (HTTP fetch + per-shard creator resolution).
|
||
///
|
||
/// Only TWO things vary between the two deployments:
|
||
///
|
||
/// 1. **partition awareness** ([`is_partitioned`](ShardCoordinator::is_partitioned)):
|
||
/// a known-down shard is pre-filtered so no worker is spawned for it.
|
||
/// 2. **per-shard creator resolution**
|
||
/// ([`resolve_creator`](ShardCoordinator::resolve_creator)): coordinator-level
|
||
/// diversity reads each item's `creator_id` from the shard that returned it.
|
||
/// In-process this is a local `get_item_metadata`; multi-process the item is
|
||
/// LOCAL to the executing region (the only place `/sharded/*` runs the merge
|
||
/// is the gateway over its own replica), so it is still a local read.
|
||
///
|
||
/// The per-shard query itself is NOT on this trait — it is the closure passed to
|
||
/// [`dispatch_shards`], which the in-process and HTTP paths each supply (a local
|
||
/// `db.retrieve` vs. a `GET {peer}/feed`). The trait deliberately stays tiny so
|
||
/// the in-process path is provably behavior-neutral (every existing
|
||
/// scatter-gather test exercises [`SimCoordinator`] unchanged).
|
||
pub trait ShardCoordinator: Sync {
|
||
/// Whether `shard` is currently known-partitioned (skip it, mark degraded).
|
||
fn is_partitioned(&self, shard: RegionId) -> bool;
|
||
|
||
/// Resolve the `creator_id` for `entity` as stored on `shard`, for
|
||
/// coordinator-level `max_per_creator` enforcement. `None` ⇒ unattributed
|
||
/// (never capped) — a metadata gap can only ever UNDER-enforce, never hide a
|
||
/// result.
|
||
fn resolve_creator(&self, shard: RegionId, entity: EntityId) -> Option<u64>;
|
||
}
|
||
|
||
/// In-process [`ShardCoordinator`] over a [`SimulatedCluster`].
|
||
///
|
||
/// Partition state comes from the cluster's partition set; creator metadata from
|
||
/// the owning node's local store. This is the EXACT behavior the pre-seam code
|
||
/// had inline, lifted behind the trait with zero change.
|
||
pub struct SimCoordinator<'a> {
|
||
cluster: &'a SimulatedCluster,
|
||
}
|
||
|
||
impl<'a> SimCoordinator<'a> {
|
||
#[must_use]
|
||
pub const fn new(cluster: &'a SimulatedCluster) -> Self {
|
||
Self { cluster }
|
||
}
|
||
}
|
||
|
||
impl ShardCoordinator for SimCoordinator<'_> {
|
||
fn is_partitioned(&self, shard: RegionId) -> bool {
|
||
self.cluster.is_partitioned(shard)
|
||
}
|
||
|
||
fn resolve_creator(&self, shard: RegionId, entity: EntityId) -> Option<u64> {
|
||
self.cluster
|
||
.node(shard)
|
||
.db
|
||
.get_item_metadata(entity)
|
||
.ok()
|
||
.flatten()
|
||
.and_then(|meta| meta.get("creator_id").and_then(|c| c.parse::<u64>().ok()))
|
||
}
|
||
}
|
||
|
||
/// Default query deadline (50ms as per spec Section 7.4).
|
||
const DEFAULT_DEADLINE_MS: u64 = 50;
|
||
|
||
/// Estimated network overhead per shard hop (subtracted from deadline).
|
||
const NETWORK_OVERHEAD_MS: u64 = 5;
|
||
|
||
/// Server-side ceiling on a client-supplied `deadline_ms`.
|
||
///
|
||
/// The scatter-gather budget is fully client-controlled (the `?deadline_ms=`
|
||
/// query param flows straight into [`scatter_gather_retrieve`] /
|
||
/// [`scatter_gather_search`]). An unbounded value lets a single request pin a
|
||
/// blocking-pool worker for an arbitrarily long time — a trivial resource-
|
||
/// exhaustion / slow-loris vector against the whole node. We clamp every
|
||
/// request to this ceiling (10s), which is already two orders of magnitude
|
||
/// above the 50ms spec target, so it never constrains a legitimate query while
|
||
/// capping the worst case. See [`clamp_deadline_ms`].
|
||
const MAX_DEADLINE_MS: u64 = 10_000;
|
||
|
||
/// Clamp a client-supplied scatter-gather deadline to [`MAX_DEADLINE_MS`].
|
||
///
|
||
/// `None` keeps the [`DEFAULT_DEADLINE_MS`] default. Any explicit value above
|
||
/// the ceiling is logged once and reduced, so a malicious or buggy client
|
||
/// cannot hold a worker indefinitely.
|
||
fn clamp_deadline_ms(requested: Option<u64>) -> u64 {
|
||
match requested {
|
||
None => DEFAULT_DEADLINE_MS,
|
||
Some(ms) if ms > MAX_DEADLINE_MS => {
|
||
tracing::warn!(
|
||
requested_ms = ms,
|
||
cap_ms = MAX_DEADLINE_MS,
|
||
"scatter-gather deadline_ms exceeds server cap; clamping"
|
||
);
|
||
MAX_DEADLINE_MS
|
||
}
|
||
Some(ms) => ms,
|
||
}
|
||
}
|
||
|
||
// ── Global scatter-gather worker pool ────────────────────────────────────────
|
||
|
||
/// Floor on the pool's persistent shard-worker threads. Even on a single-core
|
||
/// host the fan-out gets meaningful parallelism.
|
||
const MIN_SHARD_WORKERS: usize = 8;
|
||
/// Multiplier applied to available parallelism to size the pool. Shard workers
|
||
/// are query/IO-bound (a blocking `TidalDb` read), not purely CPU-bound, so a
|
||
/// modest oversubscription keeps cores busy without unbounded growth.
|
||
const SHARD_WORKERS_PER_CORE: usize = 8;
|
||
/// Hard ceiling on the pool's worker threads, independent of core count, so a
|
||
/// many-core host still keeps a bounded thread set under a query storm.
|
||
const MAX_SHARD_WORKERS: usize = 256;
|
||
/// Queued (not-yet-started) shard jobs permitted per worker before a further
|
||
/// submission is refused and the shard is reported degraded. A small multiple
|
||
/// lets a fan-out burst queue briefly rather than shed, while a sustained storm
|
||
/// still sheds promptly — the queue can never grow without bound.
|
||
const SHARD_QUEUE_DEPTH_PER_WORKER: usize = 8;
|
||
|
||
/// A boxed shard-query job run on a pool worker.
|
||
///
|
||
/// It owns everything it needs (an `Arc<Ctx>` clone, the cloned per-shard query
|
||
/// closure, the per-request result `Sender`) and sends its outcome back over
|
||
/// that channel when it finishes — or early-returns if the request budget
|
||
/// already elapsed before a worker picked it up.
|
||
type ShardJob = Box<dyn FnOnce() + Send + 'static>;
|
||
|
||
/// A fixed, process-global pool of persistent OS threads that run scatter-gather
|
||
/// shard queries.
|
||
///
|
||
/// Replaces the previous design of spawning one detached `std::thread` PER SHARD
|
||
/// PER QUERY (plus a separate counting semaphore to cap the aggregate). At high
|
||
/// read QPS across many shards, per-query thread creation/teardown dominated; a
|
||
/// pool of reused threads removes that churn entirely. The pool's fixed worker
|
||
/// count IS the aggregate concurrency cap — no separate semaphore is needed: at
|
||
/// most `workers` shard queries run at once and excess jobs queue.
|
||
///
|
||
/// The detach-on-deadline contract is preserved EXACTLY. The coordinator submits
|
||
/// one job per live shard, each carrying a clone of the per-request result
|
||
/// `Sender`, then drains the matching `Receiver` with `recv_timeout` against the
|
||
/// remaining budget. A job that finishes after the coordinator has already
|
||
/// returned sends into a `Receiver` that has been dropped; the send fails
|
||
/// harmlessly and the worker moves straight on to the next job. The coordinator
|
||
/// never joins a worker, so one slow shard's blocking query can never delay the
|
||
/// partial result — identical to the old detached-thread behavior, minus the
|
||
/// per-query spawn.
|
||
///
|
||
/// Submissions beyond the bounded queue are refused (`try_send` → `Full`); the
|
||
/// coordinator marks that shard degraded immediately rather than blocking or
|
||
/// growing the queue without bound — the pool analogue of the old "permit not
|
||
/// acquired before deadline → degrade" path.
|
||
struct ShardReadPool {
|
||
sender: crossbeam::channel::Sender<ShardJob>,
|
||
}
|
||
|
||
impl ShardReadPool {
|
||
/// Build the pool and start `workers` persistent threads draining a bounded
|
||
/// queue of depth `workers * SHARD_QUEUE_DEPTH_PER_WORKER`.
|
||
///
|
||
/// # Panics
|
||
///
|
||
/// Panics if a worker OS thread cannot be spawned. The pool is built once,
|
||
/// lazily, before any shard query runs, so a thread-exhausted host fails
|
||
/// loudly at first use rather than turning the same failure into a
|
||
/// per-request degrade on the hot path (the very hazard the pool removes).
|
||
fn new(workers: usize) -> Self {
|
||
let workers = workers.max(1);
|
||
let queue_depth = workers.saturating_mul(SHARD_QUEUE_DEPTH_PER_WORKER).max(1);
|
||
// Bounded so a sustained storm sheds (degrades shards) instead of growing
|
||
// the queue without limit. crossbeam's MPMC channel lets every worker
|
||
// pull from the same queue with no shared `Mutex<Receiver>`.
|
||
let (sender, receiver) = crossbeam::channel::bounded::<ShardJob>(queue_depth);
|
||
for i in 0..workers {
|
||
let rx: crossbeam::channel::Receiver<ShardJob> = receiver.clone();
|
||
std::thread::Builder::new()
|
||
.name(format!("scatter-shard-pool-{i}"))
|
||
.spawn(move || {
|
||
// `recv` blocks until a job arrives and only errors once every
|
||
// `Sender` is dropped — which never happens for the
|
||
// process-global pool, so these threads live for the process
|
||
// lifetime and are reused across every request.
|
||
while let Ok(job) = rx.recv() {
|
||
// Isolate a panicking shard query so it cannot retire this
|
||
// persistent worker. The old per-query design spawned a
|
||
// fresh thread each time, so a panic killed an ephemeral
|
||
// thread and the next query got a new one; a fixed pool has
|
||
// no such self-heal, so an unguarded panic would PERMANENTLY
|
||
// shrink the pool toward an all-shards-degraded wedge. The
|
||
// job sends its outcome before returning, so a caught panic
|
||
// just means that shard never reported, and the coordinator
|
||
// degrades it — the correct, honest signal.
|
||
if std::panic::catch_unwind(std::panic::AssertUnwindSafe(job)).is_err() {
|
||
tracing::error!(
|
||
worker = i,
|
||
"scatter-gather shard job panicked; worker survives, shard degraded"
|
||
);
|
||
}
|
||
}
|
||
})
|
||
.expect("spawn scatter-gather shard-pool worker thread");
|
||
}
|
||
Self { sender }
|
||
}
|
||
|
||
/// Submit a shard job. Returns `false` if the bounded queue is full (the
|
||
/// caller degrades that shard); never blocks.
|
||
fn submit(&self, job: ShardJob) -> bool {
|
||
// `Ok` ⇒ enqueued. `Full` ⇒ queue saturated → degrade. `Disconnected`
|
||
// cannot occur (the pool holds the receiver for the process lifetime),
|
||
// but treating any non-`Ok` as "not accepted" is the safe default.
|
||
self.sender.try_send(job).is_ok()
|
||
}
|
||
}
|
||
|
||
/// The process-global shard-read pool, sized once from available parallelism on
|
||
/// first use.
|
||
static SHARD_READ_POOL: OnceLock<ShardReadPool> = OnceLock::new();
|
||
|
||
/// Resolve the global shard-read pool, initializing it on first use.
|
||
fn shard_read_pool() -> &'static ShardReadPool {
|
||
SHARD_READ_POOL.get_or_init(|| {
|
||
let cores = std::thread::available_parallelism().map_or(1, std::num::NonZeroUsize::get);
|
||
let workers = cores
|
||
.saturating_mul(SHARD_WORKERS_PER_CORE)
|
||
.clamp(MIN_SHARD_WORKERS, MAX_SHARD_WORKERS);
|
||
tracing::info!(workers, "scatter-gather shard-read pool initialized");
|
||
ShardReadPool::new(workers)
|
||
})
|
||
}
|
||
|
||
/// Metadata about scatter-gather query execution.
|
||
#[derive(Debug, Clone)]
|
||
pub struct ScatterGatherMeta {
|
||
/// Whether the result set is degraded due to shard failures.
|
||
pub degraded: bool,
|
||
/// Shards that were unavailable during the query.
|
||
pub unavailable_shards: Vec<String>,
|
||
/// Number of shards that contributed results.
|
||
pub shards_queried: usize,
|
||
/// Total wall-clock time for the scatter-gather.
|
||
pub elapsed_ms: u64,
|
||
/// Per-shard deadline that was propagated.
|
||
pub shard_deadline_ms: u64,
|
||
}
|
||
|
||
/// Determines which shard owns an entity, delegating to the engine's
|
||
/// [`ShardRouter`] so server-side write/read routing can never disagree with
|
||
/// the engine's own entity→shard mapping.
|
||
///
|
||
/// The engine [`ShardRouter::hash`] uses FNV-1a over the entity ID and returns
|
||
/// a [`ShardId`] in `0..num_shards`; we index that into `shards`, which is the
|
||
/// region/shard list in ascending order, so `ShardId(i)` maps to
|
||
/// `shards[i]`. Previously this used a divergent Knuth multiplicative hash,
|
||
/// which silently disagreed with the engine and would route the same entity to
|
||
/// different shards for writes vs. reads.
|
||
///
|
||
/// # Panics
|
||
///
|
||
/// Panics in debug mode if `shards` is empty.
|
||
#[must_use]
|
||
pub fn entity_shard(entity_id: EntityId, shards: &[RegionId]) -> RegionId {
|
||
debug_assert!(!shards.is_empty(), "entity_shard requires non-empty shards");
|
||
// `ShardRouter::hash` only fails on a zero shard count, which the empty
|
||
// guard above rules out; clamp to at least one shard so the conversion is
|
||
// infallible without an `unwrap`.
|
||
let num_shards = u16::try_from(shards.len()).unwrap_or(u16::MAX).max(1);
|
||
let shard =
|
||
ShardRouter::hash(num_shards).map_or(ShardId::SINGLE, |router| router.route(entity_id));
|
||
// ShardId(i) corresponds to shards[i]; the router guarantees i < num_shards
|
||
// for a non-empty list. Fall back to the first shard rather than indexing
|
||
// out of bounds if a caller ever violates the non-empty precondition.
|
||
shards
|
||
.get(shard.0 as usize)
|
||
.or_else(|| shards.first())
|
||
.copied()
|
||
.unwrap_or(RegionId::SINGLE)
|
||
}
|
||
|
||
/// Write an item to the owning shard only (entity-sharded mode).
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// Returns [`ServerError`] if the owning shard's write fails.
|
||
// `metadata` is always the std-hasher `HashMap` built by the cluster sim; a
|
||
// generic hasher bound would add noise with no caller benefit.
|
||
#[allow(clippy::implicit_hasher)]
|
||
pub fn sharded_write_item(
|
||
cluster: &SimulatedCluster,
|
||
entity_id: EntityId,
|
||
metadata: &HashMap<String, String>,
|
||
shards: &[RegionId],
|
||
) -> Result<()> {
|
||
let shard = entity_shard(entity_id, shards);
|
||
cluster
|
||
.node(shard)
|
||
.db
|
||
.write_item_with_metadata(entity_id, metadata)
|
||
.map(|_seq| ())
|
||
.map_err(ServerError::from)
|
||
}
|
||
|
||
/// Write an embedding to the owning shard only (entity-sharded mode).
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// Returns [`ServerError`] if the owning shard's embedding write fails.
|
||
pub fn sharded_write_embedding(
|
||
cluster: &SimulatedCluster,
|
||
entity_id: EntityId,
|
||
embedding: &[f32],
|
||
shards: &[RegionId],
|
||
) -> Result<()> {
|
||
let shard = entity_shard(entity_id, shards);
|
||
cluster
|
||
.node(shard)
|
||
.db
|
||
.write_item_embedding(entity_id, embedding)
|
||
.map(|_seq| ())
|
||
.map_err(ServerError::from)
|
||
}
|
||
|
||
/// Write a signal to the owning shard (entity-sharded mode).
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// Returns [`ServerError`] if the owning shard's signal write fails.
|
||
pub fn sharded_write_signal(
|
||
cluster: &SimulatedCluster,
|
||
signal_name: &str,
|
||
entity_id: EntityId,
|
||
weight: f64,
|
||
shards: &[RegionId],
|
||
) -> Result<()> {
|
||
let shard = entity_shard(entity_id, shards);
|
||
cluster
|
||
.node(shard)
|
||
.db
|
||
.signal(
|
||
signal_name,
|
||
entity_id,
|
||
weight,
|
||
tidaldb::schema::Timestamp::now(),
|
||
)
|
||
.map_err(ServerError::from)
|
||
}
|
||
|
||
/// Resolve a human-readable shard name, falling back to `s<id>`.
|
||
fn shard_name(region_names: &HashMap<RegionId, String>, shard: RegionId) -> String {
|
||
region_names
|
||
.get(&shard)
|
||
.cloned()
|
||
.unwrap_or_else(|| format!("s{}", shard.0))
|
||
}
|
||
|
||
/// What one shard contributed to a scatter-gather.
|
||
struct ShardOutcome<T> {
|
||
items: Vec<T>,
|
||
total_candidates: usize,
|
||
}
|
||
|
||
/// A merged item tagged with the shard ([`RegionId`]) that actually returned
|
||
/// it.
|
||
///
|
||
/// Coordinator-level diversity resolves each item's creator from `region` —
|
||
/// the node that owns the item in an entity-sharded topology — rather than from
|
||
/// a single fixed replica. In a replicated topology every shard holds every
|
||
/// entity, so `region` is simply one healthy replica that has the item; in an
|
||
/// entity-sharded topology it is the *only* node that has it. Either way the
|
||
/// creator lookup hits a node that actually stores the item. See
|
||
/// [`enforce_max_per_creator`].
|
||
struct Sourced<T> {
|
||
/// The shard that returned this item (and therefore can resolve its
|
||
/// metadata locally).
|
||
region: RegionId,
|
||
item: T,
|
||
}
|
||
|
||
/// Per-shard query state accumulated by [`dispatch_shards`].
|
||
struct GatherState<T> {
|
||
/// Each merged item paired with the shard that returned it, so the merge
|
||
/// stage can resolve creators from the owning node (entity-sharded) or a
|
||
/// holding replica (replicated). See [`Sourced`].
|
||
items: Vec<Sourced<T>>,
|
||
/// Each contributing shard's `total_candidates`, kept separately so the
|
||
/// merge stage can reconcile the count instead of blindly summing — a sum
|
||
/// double-counts the candidate universe across REPLICATED shards (every
|
||
/// shard sees every entity). See [`reconcile_total_candidates`].
|
||
per_shard_totals: Vec<usize>,
|
||
unavailable_shards: Vec<String>,
|
||
shards_queried: usize,
|
||
}
|
||
|
||
/// Fan out a per-shard query CONCURRENTLY and gather the results under a hard
|
||
/// total-time budget.
|
||
///
|
||
/// Each non-partitioned shard's `query_one` runs as a job on the process-global
|
||
/// [`shard_read_pool`], holding an owned `Arc<Ctx>` clone. The underlying
|
||
/// `TidalDb` query is a *blocking* call, so true thread-level concurrency (not
|
||
/// cooperative async) is required for one slow shard not to serialize behind
|
||
/// the others — and the coordinator must NOT join a worker, so it can return its
|
||
/// partial result the instant the budget expires without waiting on a shard
|
||
/// whose blocking query is still in flight. The coordinator drains the result
|
||
/// channel with [`mpsc::Receiver::recv_timeout`] against the remaining budget.
|
||
///
|
||
/// Shards that error, or that fail to report by the deadline, are recorded in
|
||
/// `unavailable_shards` (degraded) — they are NEVER silently dropped. A pool
|
||
/// worker that finishes after the deadline simply sends into a receiver that has
|
||
/// already been dropped; the send fails harmlessly and the worker moves on to
|
||
/// the next job. Because every worker only *reads* the shared context, a
|
||
/// still-running worker outliving the request is sound.
|
||
///
|
||
/// # Fan-out cap
|
||
///
|
||
/// The pool's fixed worker count bounds the AGGREGATE number of concurrently-
|
||
/// executing shard queries regardless of how many sharded requests fan out at
|
||
/// once — no per-query thread spawn, no separate semaphore. A shard whose job
|
||
/// cannot be enqueued (the bounded queue is saturated) is reported degraded
|
||
/// immediately — never silently dropped, and never left parked. A queued job a
|
||
/// worker only reaches after the budget elapsed skips its query and reports
|
||
/// degraded too. Together with the router's request-concurrency limit this caps
|
||
/// total worker threads under a query storm instead of growing them without
|
||
/// bound.
|
||
fn dispatch_shards<Ctx, T, F>(
|
||
ctx: &Arc<Ctx>,
|
||
is_partitioned: impl Fn(RegionId) -> bool,
|
||
shards: &[RegionId],
|
||
region_names: &HashMap<RegionId, String>,
|
||
deadline: Duration,
|
||
query_one: F,
|
||
) -> GatherState<T>
|
||
where
|
||
Ctx: Send + Sync + 'static,
|
||
T: Send + 'static,
|
||
F: Fn(&Ctx, RegionId) -> Result<ShardOutcome<T>> + Clone + Send + 'static,
|
||
{
|
||
let start = Instant::now();
|
||
let mut state = GatherState {
|
||
items: Vec::new(),
|
||
per_shard_totals: Vec::new(),
|
||
unavailable_shards: Vec::new(),
|
||
shards_queried: 0,
|
||
};
|
||
|
||
// Pre-filter partitioned shards: they are known-unavailable, so we never
|
||
// spawn a worker for them.
|
||
let live_shards: Vec<RegionId> = shards
|
||
.iter()
|
||
.copied()
|
||
.filter(|&shard| {
|
||
if is_partitioned(shard) {
|
||
state
|
||
.unavailable_shards
|
||
.push(shard_name(region_names, shard));
|
||
false
|
||
} else {
|
||
true
|
||
}
|
||
})
|
||
.collect();
|
||
|
||
if live_shards.is_empty() {
|
||
return state;
|
||
}
|
||
|
||
// `(shard, Result<ShardOutcome>)` flows back over this channel. The bound
|
||
// equals the live-shard count so no worker ever blocks on send — and so a
|
||
// late worker can always deposit its result without waiting on a receiver
|
||
// that may already be gone.
|
||
let (tx, rx) = mpsc::sync_channel::<(RegionId, Result<ShardOutcome<T>>)>(live_shards.len());
|
||
|
||
// Shards whose worker actually started. A shard whose thread fails to spawn
|
||
// is marked degraded immediately and excluded from the wait set so we never
|
||
// burn the whole deadline waiting on a result that can never arrive.
|
||
let mut dispatched: Vec<RegionId> = Vec::with_capacity(live_shards.len());
|
||
let pool = shard_read_pool();
|
||
// Absolute instant the request budget expires; captured into each job so a
|
||
// job a worker only reaches after the deadline can skip its now-pointless
|
||
// query instead of piling abandoned work onto the pool under a storm.
|
||
let job_deadline = start + deadline;
|
||
for &shard in &live_shards {
|
||
let tx = tx.clone();
|
||
let ctx = Arc::clone(ctx);
|
||
let query_one = query_one.clone();
|
||
let job: ShardJob = Box::new(move || {
|
||
// If the budget elapsed before a worker picked this job up, skip the
|
||
// blocking query and report degraded. The coordinator has very
|
||
// likely already returned, so this send lands in a dropped receiver
|
||
// and fails harmlessly.
|
||
if Instant::now() >= job_deadline {
|
||
let _ = tx.send((
|
||
shard,
|
||
Err(ServerError::Unavailable(
|
||
"scatter-gather job started after deadline".into(),
|
||
)),
|
||
));
|
||
return;
|
||
}
|
||
let outcome = query_one(&ctx, shard);
|
||
// The receiver may already have moved on after the deadline; a
|
||
// closed channel is expected and benign, so the error is dropped.
|
||
let _ = tx.send((shard, outcome));
|
||
});
|
||
if pool.submit(job) {
|
||
dispatched.push(shard);
|
||
} else {
|
||
// Bounded pool queue saturated under load: degrade this shard
|
||
// immediately rather than block or grow the queue without bound. The
|
||
// other shards still race normally.
|
||
let name = shard_name(region_names, shard);
|
||
tracing::warn!(shard = shard.0, region = %name, "scatter-gather pool queue full; marking shard degraded");
|
||
state.unavailable_shards.push(name);
|
||
}
|
||
}
|
||
// Drop the coordinator's own sender so the channel closes once every worker
|
||
// has sent (and been dropped), letting `recv_timeout` observe `Disconnected`
|
||
// instead of waiting out the full budget when all shards have reported.
|
||
drop(tx);
|
||
|
||
// Shards whose result we have folded in (success OR error). Anything in
|
||
// `dispatched` but not here when the budget expires is a timed-out shard.
|
||
let mut reported: HashSet<RegionId> = HashSet::with_capacity(dispatched.len());
|
||
|
||
while reported.len() < dispatched.len() {
|
||
let remaining = deadline.saturating_sub(start.elapsed());
|
||
if remaining.is_zero() {
|
||
break;
|
||
}
|
||
match rx.recv_timeout(remaining) {
|
||
Ok((shard, outcome)) => {
|
||
fold_outcome(&mut state, &mut reported, region_names, shard, outcome);
|
||
}
|
||
Err(mpsc::RecvTimeoutError::Timeout | mpsc::RecvTimeoutError::Disconnected) => {
|
||
break;
|
||
}
|
||
}
|
||
}
|
||
|
||
// The budget elapsed (or the loop broke). Drain anything that landed in the
|
||
// race window so a straggler that finished just in time is not mis-reported
|
||
// as a timeout. We do NOT block here — the receiver is dropped right after,
|
||
// so any still-running worker's later send fails harmlessly.
|
||
while reported.len() < dispatched.len() {
|
||
match rx.try_recv() {
|
||
Ok((shard, outcome)) => {
|
||
fold_outcome(&mut state, &mut reported, region_names, shard, outcome);
|
||
}
|
||
Err(_) => break,
|
||
}
|
||
}
|
||
|
||
// Any dispatched shard with no folded result missed the deadline. Mark it
|
||
// degraded by name — never silently truncate.
|
||
let timed_out: Vec<RegionId> = dispatched
|
||
.iter()
|
||
.copied()
|
||
.filter(|shard| !reported.contains(shard))
|
||
.collect();
|
||
if !timed_out.is_empty() {
|
||
tracing::warn!(
|
||
dropped = timed_out.len(),
|
||
"scatter-gather deadline exceeded; {} shard(s) dropped as degraded",
|
||
timed_out.len()
|
||
);
|
||
for shard in timed_out {
|
||
state
|
||
.unavailable_shards
|
||
.push(shard_name(region_names, shard));
|
||
}
|
||
}
|
||
|
||
state
|
||
}
|
||
|
||
/// Fold one shard's reported outcome into the gather state and mark it reported.
|
||
fn fold_outcome<T>(
|
||
state: &mut GatherState<T>,
|
||
reported: &mut HashSet<RegionId>,
|
||
region_names: &HashMap<RegionId, String>,
|
||
shard: RegionId,
|
||
outcome: Result<ShardOutcome<T>>,
|
||
) {
|
||
reported.insert(shard);
|
||
match outcome {
|
||
Ok(outcome) => {
|
||
state.per_shard_totals.push(outcome.total_candidates);
|
||
// Tag every item with the shard that returned it so coordinator-
|
||
// level diversity can resolve its creator from a node that actually
|
||
// stores the item (the owning shard in entity-sharded mode).
|
||
state
|
||
.items
|
||
.extend(outcome.items.into_iter().map(|item| Sourced {
|
||
region: shard,
|
||
item,
|
||
}));
|
||
state.shards_queried += 1;
|
||
}
|
||
Err(e) => {
|
||
let name = shard_name(region_names, shard);
|
||
tracing::warn!(shard = shard.0, region = %name, error = %e, "shard query failed; marking degraded");
|
||
state.unavailable_shards.push(name);
|
||
}
|
||
}
|
||
}
|
||
|
||
// ── Merge helpers: dedup + candidate-count reconciliation ────────────────────
|
||
|
||
/// A merged scatter-gather item that exposes the identity, score, and rank slot
|
||
/// the coordinator needs to dedup replicated copies, merge by score, and assign
|
||
/// 1-based ranks after the merge.
|
||
trait MergeItem {
|
||
/// The entity this result is for. Replicated shards return the SAME entity,
|
||
/// so the coordinator dedups on this.
|
||
fn entity_id(&self) -> EntityId;
|
||
/// The (already-normalized) score used for K-way merge ordering.
|
||
fn score(&self) -> f64;
|
||
/// Assign the post-merge 1-based rank, so the shared merge tail can re-rank
|
||
/// either result type without duplicating the loop per query kind.
|
||
fn set_rank(&mut self, rank: usize);
|
||
}
|
||
|
||
impl MergeItem for RetrieveResult {
|
||
fn entity_id(&self) -> EntityId {
|
||
self.entity_id
|
||
}
|
||
fn score(&self) -> f64 {
|
||
self.score
|
||
}
|
||
fn set_rank(&mut self, rank: usize) {
|
||
self.rank = rank;
|
||
}
|
||
}
|
||
|
||
impl MergeItem for tidaldb::query::search::SearchResultItem {
|
||
fn entity_id(&self) -> EntityId {
|
||
self.entity_id
|
||
}
|
||
fn score(&self) -> f64 {
|
||
self.score
|
||
}
|
||
fn set_rank(&mut self, rank: usize) {
|
||
self.rank = rank;
|
||
}
|
||
}
|
||
|
||
/// Collapse duplicate entities returned by REPLICATED shards, keeping the
|
||
/// highest-scoring copy of each entity.
|
||
///
|
||
/// In a replicated topology every shard holds every entity, so the same entity
|
||
/// can appear once per shard. Without this the merged result would list the
|
||
/// same item up to `num_shards` times. The surviving copy carries its own
|
||
/// source [`RegionId`] (the shard that returned the best-scoring copy), which a
|
||
/// later creator lookup uses — that node demonstrably holds the entity, so the
|
||
/// metadata read can never miss. Returns the deduped item vector and `true` if
|
||
/// any duplicate was collapsed (i.e. shards overlapped) — the caller uses that
|
||
/// signal to reconcile `total_candidates`.
|
||
fn dedup_by_entity<T: MergeItem>(items: Vec<Sourced<T>>) -> (Vec<Sourced<T>>, bool) {
|
||
// entity_id → index of the best-scoring copy seen so far in `out`.
|
||
let mut best: HashMap<u64, usize> = HashMap::with_capacity(items.len());
|
||
let mut out: Vec<Sourced<T>> = Vec::with_capacity(items.len());
|
||
let mut overlap = false;
|
||
for sourced in items {
|
||
let key = sourced.item.entity_id().as_u64();
|
||
if let Some(idx) = best.get(&key).copied() {
|
||
overlap = true;
|
||
if sourced.item.score() > out[idx].item.score() {
|
||
out[idx] = sourced;
|
||
}
|
||
} else {
|
||
best.insert(key, out.len());
|
||
out.push(sourced);
|
||
}
|
||
}
|
||
(out, overlap)
|
||
}
|
||
|
||
/// Reconcile the merged `total_candidates` so REPLICATED shards are not
|
||
/// counted multiple times.
|
||
///
|
||
/// - **Replicated** (`overlap_detected`, i.e. [`dedup_by_entity`] collapsed at
|
||
/// least one entity returned by two shards): every shard considered the same
|
||
/// candidate universe, so the distinct count is the LARGEST single-shard
|
||
/// total, not the sum — summing inflated it by up to `num_shards`x.
|
||
/// - **Entity-sharded** (no overlap across the returned items): each shard owns
|
||
/// a disjoint slice of entities, so the totals genuinely add up.
|
||
///
|
||
/// `deduped_len` is the post-dedup merged item count, used as a floor so the
|
||
/// reported total can never be smaller than the items actually returned.
|
||
fn reconcile_total_candidates(
|
||
per_shard_totals: &[usize],
|
||
deduped_len: usize,
|
||
overlap_detected: bool,
|
||
) -> usize {
|
||
let reconciled = if overlap_detected {
|
||
// Replicated: the candidate universe is one replica's worth.
|
||
per_shard_totals.iter().copied().max().unwrap_or(0)
|
||
} else {
|
||
// Disjoint shards: the universes add up.
|
||
per_shard_totals.iter().sum()
|
||
};
|
||
reconciled.max(deduped_len)
|
||
}
|
||
|
||
/// Enforce `max_per_creator` ACROSS the merged shard results.
|
||
///
|
||
/// Each shard enforces diversity locally, but the coordinator merges several
|
||
/// shards' top-K lists, so a creator can re-appear above the cap in the merged
|
||
/// set (replicated shards each contribute the same creator; entity-sharded
|
||
/// shards each contribute their slice of a prolific creator). This walks the
|
||
/// already-score-sorted `items` and drops any item whose creator has already
|
||
/// hit the cap, mirroring the engine's per-shard `DiversitySelector` but at the
|
||
/// coordinator level.
|
||
///
|
||
/// Creators are resolved from the `creator_id` metadata of **the shard that
|
||
/// actually returned each item** ([`Sourced::region`]), not from one fixed
|
||
/// replica. This is the fix for the entity-sharded under-enforcement bug: when
|
||
/// items live on disjoint shards, the leader does not hold the ones owned by
|
||
/// other shards, so a leader-only lookup resolved their creator to `None` and
|
||
/// silently treated a prolific creator's items as uncapped. Reading from the
|
||
/// returning shard guarantees the metadata read hits a node that stores the
|
||
/// item, in both replicated and entity-sharded topologies. Items whose creator
|
||
/// still cannot be resolved (no `creator_id`, or a metadata read error) are
|
||
/// treated as having a unique, uncapped creator — never dropped — so a genuine
|
||
/// metadata gap can only ever UNDER-enforce, never hide a result.
|
||
///
|
||
/// Returns the number of items dropped, so the caller can flag the result as
|
||
/// not fully constraint-satisfied.
|
||
fn enforce_max_per_creator<T: MergeItem>(
|
||
coordinator: &dyn ShardCoordinator,
|
||
items: &mut Vec<Sourced<T>>,
|
||
max_per_creator: usize,
|
||
) -> usize {
|
||
if max_per_creator == 0 {
|
||
// A zero cap would drop everything; treat as "no cap" rather than
|
||
// silently emptying the feed — matches the engine, which ignores a 0.
|
||
return 0;
|
||
}
|
||
|
||
let mut per_creator: HashMap<u64, usize> = HashMap::new();
|
||
let mut dropped = 0usize;
|
||
|
||
items.retain(|sourced| {
|
||
// Resolve from the node that returned this item, so an entity-sharded
|
||
// owner's metadata is always reachable.
|
||
let creator = coordinator.resolve_creator(sourced.region, sourced.item.entity_id());
|
||
|
||
// Unattributed items (None) have no creator to over-represent: keep them.
|
||
creator.is_none_or(|cid| {
|
||
let count = per_creator.entry(cid).or_insert(0);
|
||
if *count >= max_per_creator {
|
||
dropped += 1;
|
||
false
|
||
} else {
|
||
*count += 1;
|
||
true
|
||
}
|
||
})
|
||
});
|
||
|
||
dropped
|
||
}
|
||
|
||
/// The type-agnostic result of merging gathered shard items: the final ranked
|
||
/// page plus the two derived facts the caller folds into its result struct.
|
||
struct MergedResults<T> {
|
||
/// The deduped, diversity-enforced, score-sorted, truncated, re-ranked page.
|
||
items: Vec<T>,
|
||
/// Candidate universe reconciled across shards (see
|
||
/// [`reconcile_total_candidates`]) — never the raw per-shard sum.
|
||
total_candidates: usize,
|
||
/// `false` iff the coordinator-level diversity pass dropped any item.
|
||
constraints_satisfied: bool,
|
||
}
|
||
|
||
/// Shared merge + assemble tail for RETRIEVE and SEARCH.
|
||
///
|
||
/// `scatter_gather_retrieve` and `scatter_gather_search` differ only in their
|
||
/// per-shard query closure and the concrete result struct they build; the merge
|
||
/// pipeline between is identical, so it lives here once:
|
||
///
|
||
/// 1. **Dedup** replicated copies of the same entity (keep the best-scoring
|
||
/// copy), recording whether shards overlapped.
|
||
/// 2. **Sort by score descending**, with `f64::total_cmp` so a NaN score yields
|
||
/// a *total*, stable order instead of being treated as equal-to-everything
|
||
/// (`partial_cmp` returns `None` for NaN, which silently degrades the sort to
|
||
/// an unstable partial order). `total_cmp` ranks NaN deterministically (a
|
||
/// positive NaN is the greatest value, so it sorts to the front of this
|
||
/// descending compare), so a single poisoned score can never interleave with
|
||
/// or scramble the ordering of the real scores.
|
||
/// 3. **Re-enforce `max_per_creator`** across the merged set when requested,
|
||
/// resolving each item's creator from the shard that returned it.
|
||
/// 4. **Reconcile `total_candidates`** BEFORE truncation so it reflects the
|
||
/// merged universe, not the page.
|
||
/// 5. **Truncate** to `limit`, drop the source-region tags, and assign 1-based
|
||
/// ranks.
|
||
fn merge_and_assemble<T: MergeItem>(
|
||
coordinator: &dyn ShardCoordinator,
|
||
items: Vec<Sourced<T>>,
|
||
per_shard_totals: &[usize],
|
||
max_per_creator: Option<usize>,
|
||
limit: usize,
|
||
) -> MergedResults<T> {
|
||
// Dedup replicated copies of the same entity (keep the best-scoring copy).
|
||
let (mut items, overlap_detected) = dedup_by_entity(items);
|
||
|
||
// Sort the deduped set by score descending. `total_cmp` gives a total order
|
||
// even if a score is NaN, so one poisoned score cannot scramble the rest.
|
||
items.sort_by(|a, b| b.item.score().total_cmp(&a.item.score()));
|
||
|
||
// Re-enforce max-per-creator across the merged set (each shard only enforced
|
||
// over its own slice). Each item's creator is resolved from the shard that
|
||
// returned it, so an entity-sharded owner's metadata is reachable.
|
||
let mut constraints_satisfied = true;
|
||
if let Some(max_per_creator) = max_per_creator {
|
||
let dropped = enforce_max_per_creator(coordinator, &mut items, max_per_creator);
|
||
if dropped > 0 {
|
||
constraints_satisfied = false;
|
||
}
|
||
}
|
||
|
||
// Reconcile the candidate count BEFORE truncating to the page limit so it
|
||
// reflects the merged universe, not the page.
|
||
let total_candidates =
|
||
reconcile_total_candidates(per_shard_totals, items.len(), overlap_detected);
|
||
|
||
// Take top limit, drop the source-region tags, and re-rank (assign 1-based
|
||
// ranks after merge).
|
||
let limit = limit.min(items.len());
|
||
items.truncate(limit);
|
||
let mut items: Vec<T> = items.into_iter().map(|s| s.item).collect();
|
||
for (i, item) in items.iter_mut().enumerate() {
|
||
item.set_rank(i + 1);
|
||
}
|
||
|
||
MergedResults {
|
||
items,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
}
|
||
}
|
||
|
||
/// Scatter-gather RETRIEVE across all shards.
|
||
///
|
||
/// Fans out the query to each non-partitioned shard CONCURRENTLY, wrapping the
|
||
/// whole gather in a hard total-time budget (per-shard deadline + network
|
||
/// overhead). The merge then, in order:
|
||
///
|
||
/// 1. **Dedups** replicated copies of the same entity (every replica returns
|
||
/// the same entity), keeping the highest-scoring copy.
|
||
/// 2. **Reconciles `total_candidates`** so replicated shards are not counted
|
||
/// multiple times (see [`reconcile_total_candidates`]) — a plain sum
|
||
/// over-counted the candidate universe by up to `num_shards`x.
|
||
/// 3. **Re-enforces `max_per_creator`** across the merged set when the query
|
||
/// declares it (see [`enforce_max_per_creator`]) — each shard only enforces
|
||
/// diversity over its own slice, and each item's creator is resolved from
|
||
/// the shard that returned it so entity-sharded owners are reachable.
|
||
/// 4. Sorts by score descending and takes the top-K.
|
||
///
|
||
/// Shards that error or miss the deadline are reported as degraded — never
|
||
/// silently truncated.
|
||
///
|
||
/// Returns the merged results and execution metadata.
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// Returns [`ServerError`] if building the per-shard query fails. Shard-level
|
||
/// failures and timeouts are reported as degraded in the metadata, not as `Err`.
|
||
// `region_names` is always the std-hasher `HashMap` owned by the cluster sim.
|
||
#[allow(clippy::implicit_hasher)]
|
||
pub fn scatter_gather_retrieve(
|
||
cluster: &Arc<SimulatedCluster>,
|
||
query: &Retrieve,
|
||
shards: &[RegionId],
|
||
region_names: &HashMap<RegionId, String>,
|
||
deadline_ms: Option<u64>,
|
||
) -> Result<(RetrieveResults, ScatterGatherMeta)> {
|
||
let start = Instant::now();
|
||
let budget_ms = clamp_deadline_ms(deadline_ms);
|
||
let total_deadline = Duration::from_millis(budget_ms);
|
||
let shard_deadline_ms = budget_ms.saturating_sub(NETWORK_OVERHEAD_MS);
|
||
|
||
// Each detached worker holds an owned clone of the query, so it can outlive
|
||
// this stack frame if the shard's blocking read exceeds the deadline.
|
||
let shared_query = Arc::new(query.clone());
|
||
let GatherState {
|
||
items: all_items,
|
||
per_shard_totals,
|
||
unavailable_shards,
|
||
shards_queried,
|
||
} = dispatch_shards::<SimulatedCluster, RetrieveResult, _>(
|
||
cluster,
|
||
|shard| cluster.is_partitioned(shard),
|
||
shards,
|
||
region_names,
|
||
total_deadline,
|
||
move |cluster, shard| {
|
||
let result = cluster
|
||
.retrieve(shard, &shared_query)
|
||
.map_err(ServerError::from)?;
|
||
Ok(ShardOutcome {
|
||
items: result.items,
|
||
total_candidates: result.total_candidates,
|
||
})
|
||
},
|
||
);
|
||
|
||
// Dedup → NaN-safe score sort → coordinator diversity → reconcile → top-K
|
||
// re-rank: the tail shared with SEARCH (see [`merge_and_assemble`]). The
|
||
// creator resolution goes through the in-process [`SimCoordinator`] seam.
|
||
let coordinator = SimCoordinator::new(cluster);
|
||
let max_per_creator = query.diversity.as_ref().and_then(|d| d.max_per_creator);
|
||
let MergedResults {
|
||
items: all_items,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
} = merge_and_assemble(
|
||
&coordinator,
|
||
all_items,
|
||
&per_shard_totals,
|
||
max_per_creator,
|
||
query.limit,
|
||
);
|
||
let limit = all_items.len();
|
||
|
||
let elapsed = start.elapsed();
|
||
let meta = ScatterGatherMeta {
|
||
degraded: !unavailable_shards.is_empty(),
|
||
unavailable_shards,
|
||
shards_queried,
|
||
elapsed_ms: elapsed.as_millis() as u64,
|
||
shard_deadline_ms,
|
||
};
|
||
|
||
let results = RetrieveResults {
|
||
items: all_items,
|
||
next_cursor: None,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
warnings: Vec::new(),
|
||
session_snapshot: None,
|
||
degradation_level: tidaldb::load::DegradationLevel::Full,
|
||
stats: tidaldb::query::stats::QueryStats {
|
||
candidates_considered: total_candidates,
|
||
candidates_after_filter: total_candidates,
|
||
candidates_after_diversity: limit,
|
||
filters_applied: 0,
|
||
scoring_time_us: 0,
|
||
diversity_time_us: 0,
|
||
total_time_us: elapsed.as_micros() as u64,
|
||
degradation_level: 0,
|
||
profile_name: String::new(),
|
||
membership_epoch: None,
|
||
},
|
||
policy_metadata: tidaldb::query::retrieve::types::PolicyMetadata::default(),
|
||
};
|
||
|
||
Ok((results, meta))
|
||
}
|
||
|
||
/// Scatter-gather SEARCH across all shards.
|
||
///
|
||
/// Fans out the search query to each non-partitioned shard CONCURRENTLY under a
|
||
/// hard total-time budget. Each shard reloads its text index inside its own
|
||
/// worker before searching. The merge then dedups replicated copies of the
|
||
/// same entity, re-enforces `max_per_creator` across the merged set when the
|
||
/// query declares it, reconciles `total_candidates` so replicated shards are
|
||
/// not counted multiple times, then sorts by score descending and takes the
|
||
/// top-K. Shards that error or miss the deadline are reported as degraded —
|
||
/// never silently truncated.
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// Returns [`ServerError`] if building the per-shard query fails. Shard-level
|
||
/// failures and timeouts are reported as degraded in the metadata, not as `Err`.
|
||
// `region_names` is always the std-hasher `HashMap` owned by the cluster sim.
|
||
#[allow(clippy::implicit_hasher)]
|
||
pub fn scatter_gather_search(
|
||
cluster: &Arc<SimulatedCluster>,
|
||
query: &Search,
|
||
shards: &[RegionId],
|
||
region_names: &HashMap<RegionId, String>,
|
||
deadline_ms: Option<u64>,
|
||
) -> Result<(SearchResults, ScatterGatherMeta)> {
|
||
let start = Instant::now();
|
||
let budget_ms = clamp_deadline_ms(deadline_ms);
|
||
let total_deadline = Duration::from_millis(budget_ms);
|
||
let shard_deadline_ms = budget_ms.saturating_sub(NETWORK_OVERHEAD_MS);
|
||
|
||
// Each detached worker holds an owned clone of the query, so it can outlive
|
||
// this stack frame if the shard's blocking search exceeds the deadline.
|
||
let shared_query = Arc::new(query.clone());
|
||
let GatherState {
|
||
items: all_items,
|
||
per_shard_totals,
|
||
unavailable_shards,
|
||
shards_queried,
|
||
} = dispatch_shards::<SimulatedCluster, tidaldb::query::search::SearchResultItem, _>(
|
||
cluster,
|
||
|shard| cluster.is_partitioned(shard),
|
||
shards,
|
||
region_names,
|
||
total_deadline,
|
||
move |cluster, shard| {
|
||
// Reload text index before searching this shard.
|
||
if let Err(e) = cluster.node(shard).db.reload_text_index() {
|
||
tracing::warn!(shard = shard.0, error = %e, "failed to reload text index");
|
||
}
|
||
let result = cluster
|
||
.search(shard, &shared_query)
|
||
.map_err(ServerError::from)?;
|
||
Ok(ShardOutcome {
|
||
items: result.items,
|
||
total_candidates: result.total_candidates,
|
||
})
|
||
},
|
||
);
|
||
|
||
// Dedup → NaN-safe score sort → coordinator diversity → reconcile → top-K
|
||
// re-rank: the tail shared with RETRIEVE (see [`merge_and_assemble`]). The
|
||
// creator resolution goes through the in-process [`SimCoordinator`] seam.
|
||
let coordinator = SimCoordinator::new(cluster);
|
||
let max_per_creator = query.diversity.as_ref().and_then(|d| d.max_per_creator);
|
||
let MergedResults {
|
||
items: all_items,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
} = merge_and_assemble(
|
||
&coordinator,
|
||
all_items,
|
||
&per_shard_totals,
|
||
max_per_creator,
|
||
query.limit as usize,
|
||
);
|
||
// Preserve the prior SEARCH stat: `candidates_after_diversity` reports the
|
||
// REQUESTED page size, not the (possibly smaller) returned count.
|
||
let limit = query.limit as usize;
|
||
|
||
let elapsed = start.elapsed();
|
||
let meta = ScatterGatherMeta {
|
||
degraded: !unavailable_shards.is_empty(),
|
||
unavailable_shards,
|
||
shards_queried,
|
||
elapsed_ms: elapsed.as_millis() as u64,
|
||
shard_deadline_ms,
|
||
};
|
||
|
||
let results = SearchResults {
|
||
items: all_items,
|
||
next_cursor: None,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
warnings: Vec::new(),
|
||
session_snapshot: None,
|
||
degradation_level: tidaldb::load::DegradationLevel::Full,
|
||
stats: tidaldb::query::stats::QueryStats {
|
||
candidates_considered: total_candidates,
|
||
candidates_after_filter: total_candidates,
|
||
candidates_after_diversity: limit,
|
||
filters_applied: 0,
|
||
scoring_time_us: 0,
|
||
diversity_time_us: 0,
|
||
total_time_us: elapsed.as_micros() as u64,
|
||
degradation_level: 0,
|
||
profile_name: String::new(),
|
||
membership_epoch: None,
|
||
},
|
||
};
|
||
|
||
Ok((results, meta))
|
||
}
|
||
|
||
// ── Multi-process HTTP scatter-gather ─────────────────────────────────────────
|
||
|
||
/// Per-shard fetch + creator-resolution context for the multi-process region node.
|
||
///
|
||
/// The local region executes against its own [`TidalDb`]; remote regions are
|
||
/// fetched over a blocking HTTP `GET {peer}/feed|/search` carrying the internal
|
||
/// marker so the peer serves locally and never re-fans-out.
|
||
///
|
||
/// This is the HTTP counterpart to [`SimCoordinator`]: it drives the SAME
|
||
/// [`dispatch_shards`] → [`merge_and_assemble`] pipeline, so the merge/dedup/
|
||
/// diversity semantics (including the honest degraded contract) are byte-for-byte
|
||
/// the single-process behavior — only the per-shard fetch and creator lookup
|
||
/// change.
|
||
pub struct HttpShardContext {
|
||
/// This gateway's own region id; `shard == local` is served locally.
|
||
local: RegionId,
|
||
/// Local engine handle for the local-region fetch and creator metadata.
|
||
db: Arc<tidaldb::TidalDb>,
|
||
/// Region id → bare `host:port` HTTP address for the remote fetch.
|
||
peer_http: HashMap<RegionId, String>,
|
||
/// Blocking client (the scatter workers are detached OS threads with no
|
||
/// tokio runtime, so the per-shard fetch cannot use the async client).
|
||
client: reqwest::blocking::Client,
|
||
/// Bearer token to forward verbatim on the remote fetch (peers share one key).
|
||
auth: Option<String>,
|
||
/// m11p7 per-node internal token, set on the marked remote fetch so the
|
||
/// owner's marker guard sees a verified sibling. `None` ⇒ no cluster key
|
||
/// (marker stays hint-only).
|
||
node_token: Option<String>,
|
||
}
|
||
|
||
impl HttpShardContext {
|
||
/// Build the context. `client` should carry no global timeout — the per-shard
|
||
/// deadline is applied per request so a slow shard degrades within budget.
|
||
#[must_use]
|
||
pub const fn new(
|
||
local: RegionId,
|
||
db: Arc<tidaldb::TidalDb>,
|
||
peer_http: HashMap<RegionId, String>,
|
||
client: reqwest::blocking::Client,
|
||
auth: Option<String>,
|
||
node_token: Option<String>,
|
||
) -> Self {
|
||
Self {
|
||
local,
|
||
db,
|
||
peer_http,
|
||
client,
|
||
auth,
|
||
node_token,
|
||
}
|
||
}
|
||
}
|
||
|
||
impl ShardCoordinator for HttpShardContext {
|
||
fn is_partitioned(&self, _shard: RegionId) -> bool {
|
||
// The region node has no leader-side partition set for read fan-out: an
|
||
// unreachable shard surfaces as a fetch error inside the worker (→
|
||
// degraded), which is the honest signal. Never pre-filter here.
|
||
false
|
||
}
|
||
|
||
fn resolve_creator(&self, shard: RegionId, entity: EntityId) -> Option<u64> {
|
||
// Only the LOCAL region's items are readable from this gateway's store.
|
||
// A remote item resolves to None (unattributed → never capped), which can
|
||
// only ever UNDER-enforce the cap, never hide a result — the same
|
||
// documented metadata-gap contract the in-process path honors.
|
||
if shard != self.local {
|
||
return None;
|
||
}
|
||
self.db
|
||
.get_item_metadata(entity)
|
||
.ok()
|
||
.flatten()
|
||
.and_then(|meta| meta.get("creator_id").and_then(|c| c.parse::<u64>().ok()))
|
||
}
|
||
}
|
||
|
||
/// One remote shard's `/feed` or `/search` JSON response, parsed into the merge
|
||
/// item type. Mirrors the wire shape produced by the region node's own handlers.
|
||
#[derive(serde::Deserialize)]
|
||
struct RemoteFeedItem {
|
||
entity_id: u64,
|
||
score: f64,
|
||
}
|
||
|
||
/// Fetch one remote region's feed over blocking HTTP within `deadline`.
|
||
fn http_fetch_feed(
|
||
ctx: &HttpShardContext,
|
||
http_addr: &str,
|
||
profile: &str,
|
||
user_id: Option<u64>,
|
||
limit: usize,
|
||
deadline: Duration,
|
||
) -> Result<ShardOutcome<RetrieveResult>> {
|
||
use std::fmt::Write as _;
|
||
let mut url = format!(
|
||
"{}?profile={}&limit={}",
|
||
crate::cluster::forward::peer_url(http_addr, "/feed"),
|
||
urlencode(profile),
|
||
limit
|
||
);
|
||
if let Some(uid) = user_id {
|
||
let _ = write!(url, "&user_id={uid}");
|
||
}
|
||
let mut req = ctx.client.get(&url).timeout(deadline).header(
|
||
crate::cluster::forward::INTERNAL_MARKER,
|
||
crate::cluster::forward::INTERNAL_MARKER_VALUE,
|
||
);
|
||
if let Some(auth) = &ctx.auth {
|
||
req = req.header(axum::http::header::AUTHORIZATION, auth);
|
||
}
|
||
if let Some(token) = &ctx.node_token {
|
||
req = req.header(crate::cluster::security::NODE_TOKEN_HEADER, token);
|
||
}
|
||
let resp = req
|
||
.send()
|
||
.map_err(|e| ServerError::Unavailable(format!("remote feed fetch failed: {e}")))?;
|
||
if !resp.status().is_success() {
|
||
return Err(ServerError::Unavailable(format!(
|
||
"remote feed returned {}",
|
||
resp.status()
|
||
)));
|
||
}
|
||
let body: serde_json::Value = resp
|
||
.json()
|
||
.map_err(|e| ServerError::Unavailable(format!("remote feed decode failed: {e}")))?;
|
||
let raw = body
|
||
.get("items")
|
||
.and_then(|v| v.as_array())
|
||
.cloned()
|
||
.unwrap_or_default();
|
||
let total = body
|
||
.get("total_candidates")
|
||
.and_then(serde_json::Value::as_u64)
|
||
.unwrap_or(raw.len() as u64) as usize;
|
||
let items = raw
|
||
.into_iter()
|
||
.filter_map(|v| serde_json::from_value::<RemoteFeedItem>(v).ok())
|
||
.map(|it| RetrieveResult {
|
||
entity_id: EntityId::new(it.entity_id),
|
||
score: it.score,
|
||
rank: 0,
|
||
signals: Vec::new(),
|
||
})
|
||
.collect();
|
||
Ok(ShardOutcome {
|
||
items,
|
||
total_candidates: total,
|
||
})
|
||
}
|
||
|
||
/// Fetch one remote region's search results over blocking HTTP within `deadline`.
|
||
fn http_fetch_search(
|
||
ctx: &HttpShardContext,
|
||
http_addr: &str,
|
||
query: &str,
|
||
user_id: Option<u64>,
|
||
limit: usize,
|
||
deadline: Duration,
|
||
) -> Result<ShardOutcome<tidaldb::query::search::SearchResultItem>> {
|
||
use std::fmt::Write as _;
|
||
let mut url = format!(
|
||
"{}?query={}&limit={}",
|
||
crate::cluster::forward::peer_url(http_addr, "/search"),
|
||
urlencode(query),
|
||
limit
|
||
);
|
||
if let Some(uid) = user_id {
|
||
let _ = write!(url, "&user_id={uid}");
|
||
}
|
||
let mut req = ctx.client.get(&url).timeout(deadline).header(
|
||
crate::cluster::forward::INTERNAL_MARKER,
|
||
crate::cluster::forward::INTERNAL_MARKER_VALUE,
|
||
);
|
||
if let Some(auth) = &ctx.auth {
|
||
req = req.header(axum::http::header::AUTHORIZATION, auth);
|
||
}
|
||
if let Some(token) = &ctx.node_token {
|
||
req = req.header(crate::cluster::security::NODE_TOKEN_HEADER, token);
|
||
}
|
||
let resp = req
|
||
.send()
|
||
.map_err(|e| ServerError::Unavailable(format!("remote search fetch failed: {e}")))?;
|
||
if !resp.status().is_success() {
|
||
return Err(ServerError::Unavailable(format!(
|
||
"remote search returned {}",
|
||
resp.status()
|
||
)));
|
||
}
|
||
let body: serde_json::Value = resp
|
||
.json()
|
||
.map_err(|e| ServerError::Unavailable(format!("remote search decode failed: {e}")))?;
|
||
let raw = body
|
||
.get("items")
|
||
.and_then(|v| v.as_array())
|
||
.cloned()
|
||
.unwrap_or_default();
|
||
let total = body
|
||
.get("total_candidates")
|
||
.and_then(serde_json::Value::as_u64)
|
||
.unwrap_or(raw.len() as u64) as usize;
|
||
let items = raw
|
||
.into_iter()
|
||
.filter_map(|v| serde_json::from_value::<RemoteFeedItem>(v).ok())
|
||
.map(|it| tidaldb::query::search::SearchResultItem {
|
||
entity_id: EntityId::new(it.entity_id),
|
||
score: it.score,
|
||
rank: 0,
|
||
bm25_score: None,
|
||
semantic_score: None,
|
||
signals: Vec::new(),
|
||
metadata: None,
|
||
})
|
||
.collect();
|
||
Ok(ShardOutcome {
|
||
items,
|
||
total_candidates: total,
|
||
})
|
||
}
|
||
|
||
/// Percent-encode a query-string value (space and the reserved set), enough for
|
||
/// the `profile` / `query` params the sharded fetch passes through.
|
||
fn urlencode(s: &str) -> String {
|
||
use std::fmt::Write as _;
|
||
let mut out = String::with_capacity(s.len());
|
||
for b in s.bytes() {
|
||
match b {
|
||
b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
|
||
out.push(b as char);
|
||
}
|
||
_ => {
|
||
let _ = write!(out, "%{b:02X}");
|
||
}
|
||
}
|
||
}
|
||
out
|
||
}
|
||
|
||
/// Multi-process scatter-gather RETRIEVE across every region from the gateway.
|
||
///
|
||
/// Local region served locally, remote regions over HTTP. Preserves the
|
||
/// in-process merge/degraded semantics verbatim via the shared pipeline.
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// Never errors on a shard-level failure (those are degraded). Returns
|
||
/// [`ServerError`] only if the merge tail itself cannot run.
|
||
#[allow(clippy::implicit_hasher)]
|
||
pub fn scatter_gather_retrieve_http(
|
||
ctx: &Arc<HttpShardContext>,
|
||
query: &Retrieve,
|
||
shards: &[RegionId],
|
||
region_names: &HashMap<RegionId, String>,
|
||
deadline_ms: Option<u64>,
|
||
) -> Result<(RetrieveResults, ScatterGatherMeta)> {
|
||
let start = Instant::now();
|
||
let budget_ms = clamp_deadline_ms(deadline_ms);
|
||
let total_deadline = Duration::from_millis(budget_ms);
|
||
let shard_deadline_ms = budget_ms.saturating_sub(NETWORK_OVERHEAD_MS);
|
||
let shard_deadline = Duration::from_millis(shard_deadline_ms.max(1));
|
||
|
||
let shared_query = Arc::new(query.clone());
|
||
let limit = query.limit;
|
||
let GatherState {
|
||
items: all_items,
|
||
per_shard_totals,
|
||
unavailable_shards,
|
||
shards_queried,
|
||
} = dispatch_shards::<HttpShardContext, RetrieveResult, _>(
|
||
ctx,
|
||
|_shard| false,
|
||
shards,
|
||
region_names,
|
||
total_deadline,
|
||
move |ctx, shard| {
|
||
if shard == ctx.local {
|
||
let result = ctx.db.retrieve(&shared_query).map_err(ServerError::from)?;
|
||
return Ok(ShardOutcome {
|
||
items: result.items,
|
||
total_candidates: result.total_candidates,
|
||
});
|
||
}
|
||
let http_addr = ctx.peer_http.get(&shard).ok_or_else(|| {
|
||
ServerError::Unavailable(format!("no http addr for shard {}", shard.0))
|
||
})?;
|
||
http_fetch_feed(
|
||
ctx,
|
||
http_addr,
|
||
&shared_query.profile.name,
|
||
shared_query.for_user,
|
||
limit,
|
||
shard_deadline,
|
||
)
|
||
},
|
||
);
|
||
|
||
let coordinator = HttpShardCoordinatorRef(ctx);
|
||
let max_per_creator = query.diversity.as_ref().and_then(|d| d.max_per_creator);
|
||
let MergedResults {
|
||
items: all_items,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
} = merge_and_assemble(
|
||
&coordinator,
|
||
all_items,
|
||
&per_shard_totals,
|
||
max_per_creator,
|
||
limit,
|
||
);
|
||
let returned = all_items.len();
|
||
|
||
let elapsed = start.elapsed();
|
||
let meta = ScatterGatherMeta {
|
||
degraded: !unavailable_shards.is_empty(),
|
||
unavailable_shards,
|
||
shards_queried,
|
||
elapsed_ms: elapsed.as_millis() as u64,
|
||
shard_deadline_ms,
|
||
};
|
||
let results = RetrieveResults {
|
||
items: all_items,
|
||
next_cursor: None,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
warnings: Vec::new(),
|
||
session_snapshot: None,
|
||
degradation_level: tidaldb::load::DegradationLevel::Full,
|
||
stats: tidaldb::query::stats::QueryStats {
|
||
candidates_considered: total_candidates,
|
||
candidates_after_filter: total_candidates,
|
||
candidates_after_diversity: returned,
|
||
filters_applied: 0,
|
||
scoring_time_us: 0,
|
||
diversity_time_us: 0,
|
||
total_time_us: elapsed.as_micros() as u64,
|
||
degradation_level: 0,
|
||
profile_name: String::new(),
|
||
membership_epoch: None,
|
||
},
|
||
policy_metadata: tidaldb::query::retrieve::types::PolicyMetadata::default(),
|
||
};
|
||
Ok((results, meta))
|
||
}
|
||
|
||
/// Multi-process scatter-gather SEARCH (HTTP counterpart of
|
||
/// [`scatter_gather_search`]).
|
||
///
|
||
/// # Errors
|
||
///
|
||
/// As [`scatter_gather_retrieve_http`].
|
||
#[allow(clippy::implicit_hasher)]
|
||
pub fn scatter_gather_search_http(
|
||
ctx: &Arc<HttpShardContext>,
|
||
query: &Search,
|
||
shards: &[RegionId],
|
||
region_names: &HashMap<RegionId, String>,
|
||
deadline_ms: Option<u64>,
|
||
) -> Result<(SearchResults, ScatterGatherMeta)> {
|
||
let start = Instant::now();
|
||
let budget_ms = clamp_deadline_ms(deadline_ms);
|
||
let total_deadline = Duration::from_millis(budget_ms);
|
||
let shard_deadline_ms = budget_ms.saturating_sub(NETWORK_OVERHEAD_MS);
|
||
let shard_deadline = Duration::from_millis(shard_deadline_ms.max(1));
|
||
|
||
let shared_query = Arc::new(query.clone());
|
||
let limit = query.limit as usize;
|
||
let GatherState {
|
||
items: all_items,
|
||
per_shard_totals,
|
||
unavailable_shards,
|
||
shards_queried,
|
||
} = dispatch_shards::<HttpShardContext, tidaldb::query::search::SearchResultItem, _>(
|
||
ctx,
|
||
|_shard| false,
|
||
shards,
|
||
region_names,
|
||
total_deadline,
|
||
move |ctx, shard| {
|
||
if shard == ctx.local {
|
||
if let Err(e) = ctx.db.reload_text_index() {
|
||
tracing::warn!(shard = shard.0, error = %e, "failed to reload text index");
|
||
}
|
||
let result = ctx.db.search(&shared_query).map_err(ServerError::from)?;
|
||
return Ok(ShardOutcome {
|
||
items: result.items,
|
||
total_candidates: result.total_candidates,
|
||
});
|
||
}
|
||
let http_addr = ctx.peer_http.get(&shard).ok_or_else(|| {
|
||
ServerError::Unavailable(format!("no http addr for shard {}", shard.0))
|
||
})?;
|
||
http_fetch_search(
|
||
ctx,
|
||
http_addr,
|
||
shared_query.query_text.as_deref().unwrap_or(""),
|
||
shared_query.for_user,
|
||
limit,
|
||
shard_deadline,
|
||
)
|
||
},
|
||
);
|
||
|
||
let coordinator = HttpShardCoordinatorRef(ctx);
|
||
let max_per_creator = query.diversity.as_ref().and_then(|d| d.max_per_creator);
|
||
let MergedResults {
|
||
items: all_items,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
} = merge_and_assemble(
|
||
&coordinator,
|
||
all_items,
|
||
&per_shard_totals,
|
||
max_per_creator,
|
||
limit,
|
||
);
|
||
|
||
let elapsed = start.elapsed();
|
||
let meta = ScatterGatherMeta {
|
||
degraded: !unavailable_shards.is_empty(),
|
||
unavailable_shards,
|
||
shards_queried,
|
||
elapsed_ms: elapsed.as_millis() as u64,
|
||
shard_deadline_ms,
|
||
};
|
||
let results = SearchResults {
|
||
items: all_items,
|
||
next_cursor: None,
|
||
total_candidates,
|
||
constraints_satisfied,
|
||
warnings: Vec::new(),
|
||
session_snapshot: None,
|
||
degradation_level: tidaldb::load::DegradationLevel::Full,
|
||
stats: tidaldb::query::stats::QueryStats {
|
||
candidates_considered: total_candidates,
|
||
candidates_after_filter: total_candidates,
|
||
candidates_after_diversity: limit,
|
||
filters_applied: 0,
|
||
scoring_time_us: 0,
|
||
diversity_time_us: 0,
|
||
total_time_us: elapsed.as_micros() as u64,
|
||
degradation_level: 0,
|
||
profile_name: String::new(),
|
||
membership_epoch: None,
|
||
},
|
||
};
|
||
Ok((results, meta))
|
||
}
|
||
|
||
/// Adapter so `&Arc<HttpShardContext>` satisfies `&dyn ShardCoordinator` for the
|
||
/// merge tail (the trait is implemented on the inner type).
|
||
struct HttpShardCoordinatorRef<'a>(&'a Arc<HttpShardContext>);
|
||
|
||
impl ShardCoordinator for HttpShardCoordinatorRef<'_> {
|
||
fn is_partitioned(&self, shard: RegionId) -> bool {
|
||
self.0.is_partitioned(shard)
|
||
}
|
||
fn resolve_creator(&self, shard: RegionId, entity: EntityId) -> Option<u64> {
|
||
self.0.resolve_creator(shard, entity)
|
||
}
|
||
}
|
||
|
||
#[cfg(test)]
|
||
// Test exemptions: unwrap on known-good fixtures + loop-counter casts in
|
||
// distribution / sharding math are idiomatic here.
|
||
#[allow(
|
||
clippy::unwrap_used,
|
||
clippy::cast_precision_loss,
|
||
clippy::cast_possible_truncation
|
||
)]
|
||
mod tests {
|
||
use std::time::Duration;
|
||
|
||
use tidaldb::{
|
||
replication::shard::RegionId,
|
||
schema::{DecaySpec, EntityKind, SchemaBuilder, Window},
|
||
testing::cluster::ClusterConfig,
|
||
};
|
||
|
||
use super::*;
|
||
|
||
fn test_schema() -> tidaldb::schema::Schema {
|
||
let mut builder = SchemaBuilder::new();
|
||
let _ = builder
|
||
.signal(
|
||
"view",
|
||
EntityKind::Item,
|
||
DecaySpec::Exponential {
|
||
half_life: Duration::from_secs(7 * 24 * 3600),
|
||
},
|
||
)
|
||
.windows(&[Window::OneHour])
|
||
.velocity(false)
|
||
.add();
|
||
builder.build().unwrap()
|
||
}
|
||
|
||
fn four_region_cluster() -> (
|
||
Arc<SimulatedCluster>,
|
||
Vec<RegionId>,
|
||
HashMap<RegionId, String>,
|
||
) {
|
||
let regions = vec![RegionId(0), RegionId(1), RegionId(2), RegionId(3)];
|
||
let config = ClusterConfig {
|
||
regions: regions.clone(),
|
||
leader_region: RegionId(0),
|
||
schema: test_schema(),
|
||
profiles: Vec::new(),
|
||
transports: None,
|
||
};
|
||
let cluster = Arc::new(SimulatedCluster::build(config));
|
||
let names: HashMap<RegionId, String> = regions
|
||
.iter()
|
||
.map(|&r| (r, format!("region-{}", r.0)))
|
||
.collect();
|
||
(cluster, regions, names)
|
||
}
|
||
|
||
#[test]
|
||
fn entity_shard_distributes_evenly() {
|
||
let shards = vec![RegionId(0), RegionId(1), RegionId(2), RegionId(3)];
|
||
let mut counts = [0u32; 4];
|
||
for i in 0..4000u64 {
|
||
let shard = entity_shard(EntityId::new(i), &shards);
|
||
counts[shard.0 as usize] += 1;
|
||
}
|
||
// Each shard should get roughly 1000 of 4000 IDs (±20%).
|
||
for (idx, &c) in counts.iter().enumerate() {
|
||
assert!(c > 800 && c < 1200, "shard {idx} got {c}, expected ~1000");
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn entity_shard_deterministic() {
|
||
let shards = vec![RegionId(0), RegionId(1), RegionId(2)];
|
||
let a = entity_shard(EntityId::new(42), &shards);
|
||
let b = entity_shard(EntityId::new(42), &shards);
|
||
assert_eq!(a, b, "same entity_id must always map to same shard");
|
||
}
|
||
|
||
/// AC1: RETRIEVE across 4 shards returns correct top-K merged by score.
|
||
#[test]
|
||
fn scatter_gather_retrieve_merges_across_shards() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
|
||
// Write items with signals to each shard (replicated cluster, so all
|
||
// shards see all data — write to leader, it replicates).
|
||
for i in 1..=20u64 {
|
||
let eid = EntityId::new(i);
|
||
cluster
|
||
.write_item_with_metadata(eid, &HashMap::new())
|
||
.unwrap();
|
||
cluster.write_signal("view", eid, i as f64).unwrap();
|
||
}
|
||
|
||
let retrieve = tidaldb::query::retrieve::Retrieve::builder()
|
||
.profile("trending")
|
||
.limit(10)
|
||
.build()
|
||
.unwrap();
|
||
|
||
let (result, meta) = scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, None)
|
||
.expect("scatter-gather should succeed");
|
||
|
||
// Should return items (up to 10).
|
||
assert!(!result.items.is_empty(), "should return items");
|
||
assert!(result.items.len() <= 10, "should respect limit");
|
||
|
||
// Scores should be in descending order.
|
||
for w in result.items.windows(2) {
|
||
assert!(
|
||
w[0].score >= w[1].score,
|
||
"scores not descending: {} < {}",
|
||
w[0].score,
|
||
w[1].score
|
||
);
|
||
}
|
||
|
||
// Ranks should be 1-based sequential.
|
||
for (i, item) in result.items.iter().enumerate() {
|
||
assert_eq!(item.rank, i + 1, "rank mismatch at position {i}");
|
||
}
|
||
|
||
// No degradation.
|
||
assert!(!meta.degraded);
|
||
assert!(meta.unavailable_shards.is_empty());
|
||
assert_eq!(meta.shards_queried, 4);
|
||
}
|
||
|
||
/// AC3: One unreachable shard returns partial results with degraded=true.
|
||
#[test]
|
||
fn scatter_gather_degraded_when_shard_partitioned() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
|
||
// Write data.
|
||
for i in 1..=10u64 {
|
||
let eid = EntityId::new(i);
|
||
cluster
|
||
.write_item_with_metadata(eid, &HashMap::new())
|
||
.unwrap();
|
||
cluster.write_signal("view", eid, 1.0).unwrap();
|
||
}
|
||
|
||
// Partition region 2.
|
||
cluster.partition_region(RegionId(2));
|
||
|
||
let retrieve = tidaldb::query::retrieve::Retrieve::builder()
|
||
.profile("trending")
|
||
.limit(10)
|
||
.build()
|
||
.unwrap();
|
||
|
||
let (result, meta) = scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, None)
|
||
.expect("scatter-gather should succeed even with partitioned shard");
|
||
|
||
// Should still return results (from 3 healthy shards).
|
||
assert!(meta.degraded, "should be degraded");
|
||
assert_eq!(meta.unavailable_shards.len(), 1);
|
||
assert_eq!(meta.unavailable_shards[0], "region-2");
|
||
assert_eq!(meta.shards_queried, 3);
|
||
|
||
// Result should still contain items (not an error).
|
||
// In replicated topology, all healthy shards have all data.
|
||
assert!(!result.items.is_empty());
|
||
}
|
||
|
||
/// AC4: Deadline propagation subtracts network overhead.
|
||
#[test]
|
||
fn scatter_gather_deadline_propagation() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
|
||
let retrieve = tidaldb::query::retrieve::Retrieve::builder()
|
||
.profile("trending")
|
||
.limit(5)
|
||
.build()
|
||
.unwrap();
|
||
|
||
// With default deadline (50ms).
|
||
let (_result, meta) =
|
||
scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, None).unwrap();
|
||
assert_eq!(
|
||
meta.shard_deadline_ms,
|
||
DEFAULT_DEADLINE_MS - NETWORK_OVERHEAD_MS,
|
||
"shard deadline should be total - overhead"
|
||
);
|
||
|
||
// With custom deadline (100ms).
|
||
let (_result, meta) =
|
||
scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, Some(100)).unwrap();
|
||
assert_eq!(meta.shard_deadline_ms, 100 - NETWORK_OVERHEAD_MS);
|
||
|
||
// With very small deadline (less than overhead).
|
||
let (_result, meta) =
|
||
scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, Some(3)).unwrap();
|
||
assert_eq!(meta.shard_deadline_ms, 0, "should saturate at 0");
|
||
}
|
||
|
||
/// A client-supplied `deadline_ms` is clamped to the server-side ceiling so
|
||
/// a single request cannot pin a worker indefinitely.
|
||
#[test]
|
||
fn deadline_ms_is_clamped_to_server_cap() {
|
||
// None → default budget.
|
||
assert_eq!(clamp_deadline_ms(None), DEFAULT_DEADLINE_MS);
|
||
// Below the cap → passed through unchanged.
|
||
assert_eq!(clamp_deadline_ms(Some(250)), 250);
|
||
// Exactly at the cap → unchanged.
|
||
assert_eq!(clamp_deadline_ms(Some(MAX_DEADLINE_MS)), MAX_DEADLINE_MS);
|
||
// Above the cap → clamped down.
|
||
assert_eq!(
|
||
clamp_deadline_ms(Some(MAX_DEADLINE_MS + 1)),
|
||
MAX_DEADLINE_MS
|
||
);
|
||
assert_eq!(clamp_deadline_ms(Some(u64::MAX)), MAX_DEADLINE_MS);
|
||
}
|
||
|
||
/// An over-budget `deadline_ms` flowing through the full retrieve path is
|
||
/// clamped: the reported `shard_deadline_ms` reflects the capped budget, not
|
||
/// the (absurd) requested one.
|
||
#[test]
|
||
fn scatter_gather_retrieve_clamps_oversized_deadline() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
let retrieve = tidaldb::query::retrieve::Retrieve::builder()
|
||
.profile("trending")
|
||
.limit(5)
|
||
.build()
|
||
.unwrap();
|
||
let (_result, meta) =
|
||
scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, Some(u64::MAX)).unwrap();
|
||
assert_eq!(
|
||
meta.shard_deadline_ms,
|
||
MAX_DEADLINE_MS - NETWORK_OVERHEAD_MS,
|
||
"oversized deadline must be clamped before overhead subtraction"
|
||
);
|
||
}
|
||
|
||
/// Scatter-gather search works across shards.
|
||
#[test]
|
||
fn scatter_gather_search_merges_across_shards() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
|
||
// Write items with text metadata for search.
|
||
for i in 1..=5u64 {
|
||
let eid = EntityId::new(i);
|
||
let mut meta = HashMap::new();
|
||
meta.insert("title".to_string(), format!("jazz piano track {i}"));
|
||
cluster.write_item_with_metadata(eid, &meta).unwrap();
|
||
cluster.write_signal("view", eid, 1.0).unwrap();
|
||
}
|
||
|
||
let search_query = tidaldb::query::search::Search::builder()
|
||
.query("jazz")
|
||
.limit(5)
|
||
.build()
|
||
.unwrap();
|
||
|
||
let (result, meta) = scatter_gather_search(&cluster, &search_query, &shards, &names, None)
|
||
.expect("scatter-gather search should succeed");
|
||
|
||
// Search may or may not find results depending on text index reload timing,
|
||
// but the scatter-gather itself should not error.
|
||
assert!(!meta.degraded, "should not be degraded");
|
||
assert_eq!(meta.shards_queried, 4);
|
||
// Scores descending if items returned.
|
||
for w in result.items.windows(2) {
|
||
assert!(w[0].score >= w[1].score);
|
||
}
|
||
}
|
||
|
||
/// `SimulatedCluster` must be `Sync` for the detached scatter-gather
|
||
/// workers (each holds a shared `Arc` and only reads). A regression that
|
||
/// made it non-`Sync` would break the whole fan-out design.
|
||
#[test]
|
||
fn simulated_cluster_is_sync() {
|
||
fn assert_sync<T: Sync + Send + 'static>() {}
|
||
assert_sync::<SimulatedCluster>();
|
||
}
|
||
|
||
/// SCATTER-1: a single hung shard must NOT block the whole gather past its
|
||
/// deadline. The slow shard is reported as degraded; the fast shards still
|
||
/// contribute. Exercises [`dispatch_shards`] directly with a closure that
|
||
/// sleeps one shard well past the deadline.
|
||
#[test]
|
||
fn dispatch_shards_slow_shard_does_not_block_deadline() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
let slow_shard = shards[2];
|
||
let deadline = Duration::from_millis(60);
|
||
|
||
let started = Instant::now();
|
||
let state = dispatch_shards::<SimulatedCluster, u64, _>(
|
||
&cluster,
|
||
|shard| cluster.is_partitioned(shard),
|
||
&shards,
|
||
&names,
|
||
deadline,
|
||
move |_cluster, shard| {
|
||
if shard == slow_shard {
|
||
// Far longer than the deadline — simulates a hung shard.
|
||
std::thread::sleep(Duration::from_secs(2));
|
||
}
|
||
Ok(ShardOutcome {
|
||
items: vec![u64::from(shard.0)],
|
||
total_candidates: 1,
|
||
})
|
||
},
|
||
);
|
||
let elapsed = started.elapsed();
|
||
|
||
// The coordinator returned without waiting on the hung shard. Allow a
|
||
// generous ceiling for thread-spawn + scheduling jitter, but it must be
|
||
// far below the 2s the slow worker sleeps.
|
||
assert!(
|
||
elapsed < Duration::from_millis(800),
|
||
"gather blocked on slow shard: took {elapsed:?}"
|
||
);
|
||
|
||
// The three fast shards contributed; the slow shard is degraded, not
|
||
// silently dropped.
|
||
assert_eq!(state.shards_queried, 3, "fast shards should contribute");
|
||
assert_eq!(
|
||
state.unavailable_shards,
|
||
vec![shard_name(&names, slow_shard)],
|
||
"slow shard must be reported degraded"
|
||
);
|
||
assert_eq!(state.items.len(), 3);
|
||
}
|
||
|
||
/// SCATTER-1: with NO slow shard, all shards report and nothing is degraded
|
||
/// even under a tight-but-sufficient budget.
|
||
#[test]
|
||
fn dispatch_shards_all_report_under_budget() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
let state = dispatch_shards::<SimulatedCluster, u64, _>(
|
||
&cluster,
|
||
|shard| cluster.is_partitioned(shard),
|
||
&shards,
|
||
&names,
|
||
Duration::from_millis(500),
|
||
|_cluster, shard| {
|
||
Ok(ShardOutcome {
|
||
items: vec![u64::from(shard.0)],
|
||
total_candidates: 2,
|
||
})
|
||
},
|
||
);
|
||
assert_eq!(state.shards_queried, 4);
|
||
assert!(state.unavailable_shards.is_empty());
|
||
// Each of the 4 shards reported total_candidates == 2 (completion order
|
||
// is non-deterministic, so compare the sorted set, not the sequence).
|
||
let mut totals = state.per_shard_totals.clone();
|
||
totals.sort_unstable();
|
||
assert_eq!(totals, vec![2, 2, 2, 2]);
|
||
assert_eq!(state.per_shard_totals.iter().sum::<usize>(), 8);
|
||
assert_eq!(state.items.len(), 4);
|
||
}
|
||
|
||
/// SCATTER-1: a shard whose query errors is degraded; the rest still merge.
|
||
#[test]
|
||
fn dispatch_shards_error_shard_is_degraded() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
let bad_shard = shards[1];
|
||
let state = dispatch_shards::<SimulatedCluster, u64, _>(
|
||
&cluster,
|
||
|shard| cluster.is_partitioned(shard),
|
||
&shards,
|
||
&names,
|
||
Duration::from_millis(500),
|
||
move |_cluster, shard| {
|
||
if shard == bad_shard {
|
||
Err(ServerError::BadRequest("boom".into()))
|
||
} else {
|
||
Ok(ShardOutcome {
|
||
items: vec![u64::from(shard.0)],
|
||
total_candidates: 1,
|
||
})
|
||
}
|
||
},
|
||
);
|
||
assert_eq!(state.shards_queried, 3);
|
||
assert_eq!(
|
||
state.unavailable_shards,
|
||
vec![shard_name(&names, bad_shard)]
|
||
);
|
||
}
|
||
|
||
fn retrieve_result(entity_id: u64, score: f64) -> RetrieveResult {
|
||
RetrieveResult {
|
||
entity_id: EntityId::new(entity_id),
|
||
score,
|
||
rank: 0,
|
||
signals: Vec::new(),
|
||
}
|
||
}
|
||
|
||
/// Wrap a [`RetrieveResult`] as if it were returned by `region`, for the
|
||
/// diversity/merge unit tests that exercise [`Sourced`]-keyed paths.
|
||
fn sourced(region: RegionId, entity_id: u64, score: f64) -> Sourced<RetrieveResult> {
|
||
Sourced {
|
||
region,
|
||
item: retrieve_result(entity_id, score),
|
||
}
|
||
}
|
||
|
||
/// `entity_shard` must agree with the engine `ShardRouter::hash` (FNV-1a)
|
||
/// for EVERY entity, so server-side write/read routing can never disagree
|
||
/// with the engine's own mapping. A divergent hash silently sends the same
|
||
/// entity to different shards for writes vs reads.
|
||
#[test]
|
||
fn entity_shard_matches_engine_router() {
|
||
let shards = vec![RegionId(0), RegionId(1), RegionId(2), RegionId(3)];
|
||
let router = ShardRouter::hash(shards.len() as u16).unwrap();
|
||
for i in 0..2000u64 {
|
||
let via_server = entity_shard(EntityId::new(i), &shards);
|
||
let via_engine = router.route(EntityId::new(i));
|
||
assert_eq!(
|
||
via_server.0, via_engine.0,
|
||
"entity {i}: server routed to {via_server:?} but engine routed to {via_engine:?}"
|
||
);
|
||
}
|
||
}
|
||
|
||
/// Replicated shards return the SAME entity, so the merge must collapse
|
||
/// duplicates and keep the highest-scoring copy — and report that an
|
||
/// overlap was detected.
|
||
#[test]
|
||
fn dedup_collapses_replicated_copies_keeping_best_score() {
|
||
// Entity 1 returned by three replicas with different scores; entity 2
|
||
// by two replicas; entity 3 once. The best-scoring copy's source region
|
||
// must survive (it is the one a later creator lookup reads from).
|
||
let items = vec![
|
||
sourced(RegionId(0), 1, 0.5),
|
||
sourced(RegionId(1), 1, 0.9), // best for entity 1, from region 1
|
||
sourced(RegionId(2), 1, 0.7),
|
||
sourced(RegionId(0), 2, 0.3),
|
||
sourced(RegionId(3), 2, 0.4), // best for entity 2, from region 3
|
||
sourced(RegionId(2), 3, 0.8),
|
||
];
|
||
let (deduped, overlap) = dedup_by_entity(items);
|
||
assert!(overlap, "duplicates across replicas must set overlap=true");
|
||
assert_eq!(deduped.len(), 3, "one row per distinct entity");
|
||
|
||
let mut by_id: HashMap<u64, (f64, RegionId)> = HashMap::new();
|
||
for s in &deduped {
|
||
by_id.insert(s.item.entity_id.as_u64(), (s.item.score, s.region));
|
||
}
|
||
assert!((by_id[&1].0 - 0.9).abs() < f64::EPSILON, "kept best for e1");
|
||
assert_eq!(by_id[&1].1, RegionId(1), "kept the winning copy's region");
|
||
assert!((by_id[&2].0 - 0.4).abs() < f64::EPSILON, "kept best for e2");
|
||
assert_eq!(by_id[&2].1, RegionId(3), "kept the winning copy's region");
|
||
assert!((by_id[&3].0 - 0.8).abs() < f64::EPSILON);
|
||
}
|
||
|
||
/// Disjoint (entity-sharded) shards never return the same entity, so dedup
|
||
/// is a no-op and reports no overlap.
|
||
#[test]
|
||
fn dedup_no_overlap_for_disjoint_shards() {
|
||
let items = vec![
|
||
sourced(RegionId(0), 10, 0.5),
|
||
sourced(RegionId(1), 20, 0.6),
|
||
sourced(RegionId(2), 30, 0.7),
|
||
];
|
||
let (deduped, overlap) = dedup_by_entity(items);
|
||
assert!(!overlap, "disjoint shards must report overlap=false");
|
||
assert_eq!(deduped.len(), 3);
|
||
}
|
||
|
||
/// REPLICATED: the candidate universe is ONE replica's worth, not the sum.
|
||
/// Three replicas each reporting 100 candidates is 100 distinct, not 300.
|
||
#[test]
|
||
fn reconcile_replicated_uses_max_not_sum() {
|
||
let total = reconcile_total_candidates(&[100, 100, 100], 10, /*overlap*/ true);
|
||
assert_eq!(total, 100, "replicated shards must not be summed");
|
||
}
|
||
|
||
/// ENTITY-SHARDED: disjoint universes genuinely add up.
|
||
#[test]
|
||
fn reconcile_disjoint_sums() {
|
||
let total = reconcile_total_candidates(&[30, 40, 30], 50, /*overlap*/ false);
|
||
assert_eq!(total, 100, "disjoint shards must be summed");
|
||
}
|
||
|
||
/// The reported total can never be smaller than the items actually
|
||
/// returned (defends against a stale/under-reported per-shard total).
|
||
#[test]
|
||
fn reconcile_floors_at_deduped_len() {
|
||
let total = reconcile_total_candidates(&[2, 2], 5, /*overlap*/ true);
|
||
assert_eq!(total, 5, "must floor at the deduped item count");
|
||
}
|
||
|
||
/// A single NaN score must NOT scramble the ordering of the real-scored
|
||
/// items. `partial_cmp(...).unwrap_or(Equal)` treated NaN as equal to
|
||
/// everything, degrading the sort to an unstable partial order where one
|
||
/// poisoned candidate could reorder the rest; `total_cmp` gives a total
|
||
/// order so the real scores stay strictly descending and the NaN item is
|
||
/// ranked deterministically — never dropped, and never able to interleave
|
||
/// with the real scores.
|
||
#[test]
|
||
fn merge_and_assemble_nan_score_does_not_scramble_order() {
|
||
let (cluster, _shards, _names) = four_region_cluster();
|
||
let leader = cluster.leader_region();
|
||
// No creator_id → diversity is a no-op and every item is retained, so we
|
||
// observe pure sort behavior. Entity 3 carries a NaN score.
|
||
let items = vec![
|
||
sourced(leader, 1, 0.2),
|
||
sourced(leader, 2, 0.9),
|
||
sourced(leader, 3, f64::NAN),
|
||
sourced(leader, 4, 0.5),
|
||
];
|
||
let merged = merge_and_assemble(&SimCoordinator::new(&cluster), items, &[4], None, 10);
|
||
|
||
// Nothing dropped (no diversity cap, limit exceeds the set).
|
||
assert_eq!(
|
||
merged.items.len(),
|
||
4,
|
||
"NaN must not drop or duplicate items"
|
||
);
|
||
assert!(merged.constraints_satisfied);
|
||
|
||
// The three real-scored items stay strictly descending regardless of
|
||
// where the NaN landed.
|
||
let real: Vec<f64> = merged
|
||
.items
|
||
.iter()
|
||
.map(MergeItem::score)
|
||
.filter(|s| !s.is_nan())
|
||
.collect();
|
||
assert_eq!(
|
||
real,
|
||
vec![0.9, 0.5, 0.2],
|
||
"real scores must stay descending"
|
||
);
|
||
|
||
// `total_cmp` ranks a (positive) NaN as the greatest value, so in this
|
||
// descending compare it sorts to the FRONT — deterministically, every
|
||
// run — rather than randomly interleaving among the real scores as the
|
||
// old partial_cmp-as-Equal path allowed.
|
||
assert!(
|
||
merged.items.first().is_some_and(|s| s.score().is_nan()),
|
||
"NaN item must sort deterministically, not scramble the rest"
|
||
);
|
||
|
||
// Exactly one NaN survives and it is the only non-finite entry.
|
||
let nan_count = merged.items.iter().filter(|s| s.score().is_nan()).count();
|
||
assert_eq!(nan_count, 1, "the single NaN item is retained exactly once");
|
||
|
||
// Ranks are 1-based and contiguous over the whole merged page.
|
||
for (i, item) in merged.items.iter().enumerate() {
|
||
assert_eq!(item.rank, i + 1, "rank must be 1-based contiguous");
|
||
}
|
||
}
|
||
|
||
/// The shared merge tail honors the page `limit` and assigns 1-based ranks
|
||
/// for SEARCH items too (proving [`merge_and_assemble`] is generic over the
|
||
/// result type, not just `RetrieveResult`).
|
||
#[test]
|
||
fn merge_and_assemble_truncates_and_ranks_search_items() {
|
||
use tidaldb::query::search::SearchResultItem;
|
||
|
||
let (cluster, _shards, _names) = four_region_cluster();
|
||
let leader = cluster.leader_region();
|
||
let make = |entity_id: u64, score: f64| Sourced {
|
||
region: leader,
|
||
item: SearchResultItem {
|
||
entity_id: EntityId::new(entity_id),
|
||
score,
|
||
rank: 0,
|
||
bm25_score: None,
|
||
semantic_score: None,
|
||
signals: Vec::new(),
|
||
metadata: None,
|
||
},
|
||
};
|
||
let items = vec![make(1, 0.1), make(2, 0.9), make(3, 0.5), make(4, 0.7)];
|
||
let merged = merge_and_assemble(&SimCoordinator::new(&cluster), items, &[4], None, 2);
|
||
|
||
assert_eq!(merged.items.len(), 2, "limit must truncate to the page");
|
||
// Top-2 by score descending.
|
||
assert!((merged.items[0].score - 0.9).abs() < f64::EPSILON);
|
||
assert!((merged.items[1].score - 0.7).abs() < f64::EPSILON);
|
||
assert_eq!(merged.items[0].rank, 1);
|
||
assert_eq!(merged.items[1].rank, 2);
|
||
}
|
||
|
||
/// End-to-end: across a 4-way REPLICATED cluster, `total_candidates` must
|
||
/// NOT be 4x the single-shard count, and the merged items must not contain
|
||
/// duplicate entities.
|
||
#[test]
|
||
fn scatter_gather_retrieve_total_candidates_not_double_counted() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
for i in 1..=12u64 {
|
||
let eid = EntityId::new(i);
|
||
cluster
|
||
.write_item_with_metadata(eid, &HashMap::new())
|
||
.unwrap();
|
||
cluster.write_signal("view", eid, i as f64).unwrap();
|
||
}
|
||
|
||
// Single-shard baseline candidate count.
|
||
let retrieve = tidaldb::query::retrieve::Retrieve::builder()
|
||
.profile("trending")
|
||
.limit(12)
|
||
.build()
|
||
.unwrap();
|
||
let single = cluster.retrieve(shards[0], &retrieve).unwrap();
|
||
let single_total = single.total_candidates;
|
||
assert!(single_total > 0, "baseline shard should see candidates");
|
||
|
||
let (result, meta) =
|
||
scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, Some(500)).unwrap();
|
||
assert_eq!(meta.shards_queried, 4);
|
||
|
||
// Replicated: merged total must equal one replica's universe, NOT 4x.
|
||
assert_eq!(
|
||
result.total_candidates, single_total,
|
||
"replicated shards double-counted: {} != {single_total}",
|
||
result.total_candidates
|
||
);
|
||
|
||
// No duplicate entities survived the merge.
|
||
let mut seen = HashSet::new();
|
||
for item in &result.items {
|
||
assert!(
|
||
seen.insert(item.entity_id.as_u64()),
|
||
"duplicate entity {} in merged result",
|
||
item.entity_id.as_u64()
|
||
);
|
||
}
|
||
}
|
||
|
||
/// End-to-end: coordinator-level `max_per_creator` must hold across the
|
||
/// MERGED set, not just per shard. Twelve items all from creator 7; with
|
||
/// `max_per_creator = 3` the merged feed may contain at most 3.
|
||
#[test]
|
||
fn scatter_gather_retrieve_enforces_coordinator_diversity() {
|
||
let (cluster, shards, names) = four_region_cluster();
|
||
for i in 1..=12u64 {
|
||
let eid = EntityId::new(i);
|
||
let mut meta = HashMap::new();
|
||
meta.insert("creator_id".to_string(), "7".to_string());
|
||
cluster.write_item_with_metadata(eid, &meta).unwrap();
|
||
cluster.write_signal("view", eid, i as f64).unwrap();
|
||
}
|
||
|
||
let retrieve = tidaldb::query::retrieve::Retrieve::builder()
|
||
.profile("trending")
|
||
.limit(12)
|
||
.diversity(tidaldb::ranking::diversity::DiversityConstraints::new().max_per_creator(3))
|
||
.build()
|
||
.unwrap();
|
||
|
||
let (result, _meta) =
|
||
scatter_gather_retrieve(&cluster, &retrieve, &shards, &names, Some(500)).unwrap();
|
||
|
||
let from_creator_7 = result.items.len();
|
||
assert!(
|
||
from_creator_7 <= 3,
|
||
"coordinator diversity breached: {from_creator_7} items from one creator (cap 3)"
|
||
);
|
||
assert!(
|
||
!result.constraints_satisfied,
|
||
"dropping items for the cap must clear constraints_satisfied"
|
||
);
|
||
}
|
||
|
||
/// A zero `max_per_creator` is treated as "no cap" (matches the engine),
|
||
/// never as "drop everything".
|
||
#[test]
|
||
fn enforce_max_per_creator_zero_is_no_cap() {
|
||
let (cluster, _shards, _names) = four_region_cluster();
|
||
let leader = cluster.leader_region();
|
||
let mut meta = HashMap::new();
|
||
meta.insert("creator_id".to_string(), "7".to_string());
|
||
for i in 1..=3u64 {
|
||
cluster
|
||
.write_item_with_metadata(EntityId::new(i), &meta)
|
||
.unwrap();
|
||
}
|
||
let mut items = vec![
|
||
sourced(leader, 1, 0.9),
|
||
sourced(leader, 2, 0.8),
|
||
sourced(leader, 3, 0.7),
|
||
];
|
||
let dropped = enforce_max_per_creator(&SimCoordinator::new(&cluster), &mut items, 0);
|
||
assert_eq!(dropped, 0, "zero cap must not drop anything");
|
||
assert_eq!(items.len(), 3);
|
||
}
|
||
|
||
/// Items with no resolvable creator are never dropped by the cap (a missing
|
||
/// `creator_id` can only ever UNDER-enforce, never hide a result).
|
||
#[test]
|
||
fn enforce_max_per_creator_keeps_unattributed_items() {
|
||
let (cluster, _shards, _names) = four_region_cluster();
|
||
let leader = cluster.leader_region();
|
||
// Write items WITHOUT creator_id.
|
||
for i in 1..=5u64 {
|
||
cluster
|
||
.write_item_with_metadata(EntityId::new(i), &HashMap::new())
|
||
.unwrap();
|
||
}
|
||
let mut items: Vec<Sourced<RetrieveResult>> = (1..=5u64)
|
||
.map(|i| sourced(leader, i, 1.0 / i as f64))
|
||
.collect();
|
||
let dropped = enforce_max_per_creator(&SimCoordinator::new(&cluster), &mut items, 1);
|
||
assert_eq!(dropped, 0, "unattributed items must never be capped");
|
||
assert_eq!(items.len(), 5);
|
||
}
|
||
|
||
/// Entity-sharded under-enforcement regression: items owned by different
|
||
/// shards must be capped by resolving each item's creator from the shard
|
||
/// that returned it. Writing the same prolific creator's items to disjoint
|
||
/// shards (so the leader does NOT hold the non-leader ones) and merging them
|
||
/// must still respect the cap — the old leader-only lookup resolved the
|
||
/// non-leader items' creator to `None` and let them through uncapped.
|
||
#[test]
|
||
fn enforce_max_per_creator_resolves_from_owning_shard() {
|
||
let (cluster, _shards, _names) = four_region_cluster();
|
||
let regions = [RegionId(0), RegionId(1), RegionId(2), RegionId(3)];
|
||
|
||
// Write 8 items from creator 7, each to a DIFFERENT shard's local store
|
||
// only (no replication) — exactly the entity-sharded layout where the
|
||
// leader holds only its own slice.
|
||
let mut meta = HashMap::new();
|
||
meta.insert("creator_id".to_string(), "7".to_string());
|
||
let mut items: Vec<Sourced<RetrieveResult>> = Vec::new();
|
||
for i in 1..=8u64 {
|
||
let region = regions[(i as usize - 1) % regions.len()];
|
||
cluster
|
||
.node(region)
|
||
.db
|
||
.write_item_with_metadata(EntityId::new(i), &meta)
|
||
.unwrap();
|
||
// Source-tag each item with the shard that "returned" it.
|
||
items.push(sourced(region, i, 1.0 / i as f64));
|
||
}
|
||
|
||
let dropped = enforce_max_per_creator(&SimCoordinator::new(&cluster), &mut items, 3);
|
||
assert_eq!(
|
||
items.len(),
|
||
3,
|
||
"creator cap must hold across shards; got {} items",
|
||
items.len()
|
||
);
|
||
assert_eq!(dropped, 5, "8 items from one creator, cap 3 → drop 5");
|
||
}
|
||
|
||
/// Counter-test proving the bug the fix removes: resolving every item's
|
||
/// creator from the LEADER only (the old behavior) under-enforces, because
|
||
/// the leader does not hold the non-leader shards' items.
|
||
#[test]
|
||
fn leader_only_creator_lookup_under_enforces_when_sharded() {
|
||
let (cluster, _shards, _names) = four_region_cluster();
|
||
let regions = [RegionId(0), RegionId(1), RegionId(2), RegionId(3)];
|
||
let leader = cluster.leader_region();
|
||
|
||
let mut meta = HashMap::new();
|
||
meta.insert("creator_id".to_string(), "7".to_string());
|
||
let mut leader_only: HashMap<u64, usize> = HashMap::new();
|
||
let mut kept = 0usize;
|
||
for i in 1..=8u64 {
|
||
let region = regions[(i as usize - 1) % regions.len()];
|
||
cluster
|
||
.node(region)
|
||
.db
|
||
.write_item_with_metadata(EntityId::new(i), &meta)
|
||
.unwrap();
|
||
// Old behavior: always read metadata from the leader.
|
||
let creator = cluster
|
||
.node(leader)
|
||
.db
|
||
.get_item_metadata(EntityId::new(i))
|
||
.ok()
|
||
.flatten()
|
||
.and_then(|m| m.get("creator_id").and_then(|c| c.parse::<u64>().ok()));
|
||
let pass = creator.is_none_or(|cid| {
|
||
let c = leader_only.entry(cid).or_insert(0);
|
||
if *c >= 3 {
|
||
false
|
||
} else {
|
||
*c += 1;
|
||
true
|
||
}
|
||
});
|
||
if pass {
|
||
kept += 1;
|
||
}
|
||
}
|
||
// The leader holds only its own ~2 items, so the other ~6 resolve to
|
||
// None and slip through uncapped — proving the silent under-enforcement
|
||
// the per-shard lookup fixes.
|
||
assert!(
|
||
kept > 3,
|
||
"leader-only lookup should under-enforce (kept {kept} > cap 3)"
|
||
);
|
||
}
|
||
|
||
/// C16: the shard-read pool caps the AGGREGATE concurrent shard workers at
|
||
/// its fixed worker count. With 2 workers, no more than 2 jobs run at once
|
||
/// even when 8 are submitted, and every job eventually completes (no
|
||
/// deadlock, threads reused across the burst).
|
||
#[test]
|
||
fn shard_read_pool_caps_concurrency() {
|
||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||
|
||
let pool = ShardReadPool::new(2);
|
||
let live = Arc::new(AtomicUsize::new(0));
|
||
let peak = Arc::new(AtomicUsize::new(0));
|
||
let completed = Arc::new(AtomicUsize::new(0));
|
||
|
||
for _ in 0..8 {
|
||
let live = Arc::clone(&live);
|
||
let peak = Arc::clone(&peak);
|
||
let completed = Arc::clone(&completed);
|
||
let job: ShardJob = Box::new(move || {
|
||
let now = live.fetch_add(1, Ordering::SeqCst) + 1;
|
||
peak.fetch_max(now, Ordering::SeqCst);
|
||
// Hold the worker briefly so contention is real.
|
||
std::thread::sleep(Duration::from_millis(5));
|
||
live.fetch_sub(1, Ordering::SeqCst);
|
||
completed.fetch_add(1, Ordering::SeqCst);
|
||
});
|
||
assert!(pool.submit(job), "queue (depth 16) accepts all 8 jobs");
|
||
}
|
||
|
||
// Spin-wait for all jobs to finish (generous ceiling; no fixed sleep).
|
||
let start = Instant::now();
|
||
while completed.load(Ordering::SeqCst) < 8 {
|
||
assert!(
|
||
start.elapsed() < Duration::from_secs(5),
|
||
"every job must complete (no deadlock); saw {}",
|
||
completed.load(Ordering::SeqCst)
|
||
);
|
||
std::thread::yield_now();
|
||
}
|
||
|
||
assert!(
|
||
peak.load(Ordering::SeqCst) <= 2,
|
||
"no more than 2 jobs may run at once, saw {}",
|
||
peak.load(Ordering::SeqCst)
|
||
);
|
||
assert_eq!(
|
||
completed.load(Ordering::SeqCst),
|
||
8,
|
||
"every job must complete (no deadlock, workers reused)"
|
||
);
|
||
}
|
||
|
||
/// C16: a saturated pool refuses further submissions (so the coordinator
|
||
/// degrades that shard) instead of blocking or growing the queue. One
|
||
/// worker, queue depth 8: occupy the worker plus fill the 8 queue slots,
|
||
/// then prove the next submit is refused.
|
||
#[test]
|
||
fn shard_read_pool_refuses_when_saturated() {
|
||
let pool = ShardReadPool::new(1);
|
||
|
||
// Park the sole worker until released, so nothing drains the queue.
|
||
let (started_tx, started_rx) = std::sync::mpsc::channel::<()>();
|
||
let (release_tx, release_rx) = std::sync::mpsc::channel::<()>();
|
||
let park: ShardJob = Box::new(move || {
|
||
started_tx.send(()).expect("signal worker started");
|
||
release_rx.recv().expect("wait for release");
|
||
});
|
||
assert!(pool.submit(park), "worker accepts the parking job");
|
||
started_rx
|
||
.recv_timeout(Duration::from_secs(5))
|
||
.expect("worker started the parking job");
|
||
|
||
// Fill the 8 queue slots (the worker is busy on `park`).
|
||
for _ in 0..SHARD_QUEUE_DEPTH_PER_WORKER {
|
||
let filler: ShardJob = Box::new(|| {});
|
||
assert!(pool.submit(filler), "queue slot accepts a filler job");
|
||
}
|
||
|
||
// Queue full: the next submit is refused so the caller degrades.
|
||
let overflow: ShardJob = Box::new(|| {});
|
||
assert!(
|
||
!pool.submit(overflow),
|
||
"saturated pool must refuse so the shard is degraded, not block"
|
||
);
|
||
|
||
// Release the worker so the pool drains cleanly.
|
||
release_tx.send(()).expect("release the worker");
|
||
}
|
||
|
||
/// A panicking shard job must NOT retire its worker — the persistent pool has
|
||
/// no per-query self-heal, so an unguarded panic would permanently shrink it.
|
||
/// One worker: panic a job, then prove a later job still runs (worker alive).
|
||
#[test]
|
||
fn shard_read_pool_survives_panicking_job() {
|
||
let pool = ShardReadPool::new(1);
|
||
|
||
// Silence the default panic hook for the duration so the deliberately
|
||
// panicking job does not spam the test output with a backtrace.
|
||
let prev = std::panic::take_hook();
|
||
std::panic::set_hook(Box::new(|_| {}));
|
||
|
||
let (panicked_tx, panicked_rx) = std::sync::mpsc::channel::<()>();
|
||
let panic_job: ShardJob = Box::new(move || {
|
||
// Signal BEFORE panicking so the test knows the worker reached it.
|
||
panicked_tx.send(()).expect("signal pre-panic");
|
||
panic!("deliberate shard-job panic");
|
||
});
|
||
assert!(pool.submit(panic_job), "pool accepts the panicking job");
|
||
panicked_rx
|
||
.recv_timeout(Duration::from_secs(5))
|
||
.expect("worker ran the panicking job");
|
||
|
||
std::panic::set_hook(prev);
|
||
|
||
// The SAME (sole) worker must still process the next job — if the panic
|
||
// had killed it, this would never run and the recv would time out.
|
||
let (done_tx, done_rx) = std::sync::mpsc::channel::<()>();
|
||
let after: ShardJob = Box::new(move || {
|
||
done_tx.send(()).expect("signal post-panic job ran");
|
||
});
|
||
assert!(pool.submit(after), "pool accepts the post-panic job");
|
||
done_rx
|
||
.recv_timeout(Duration::from_secs(5))
|
||
.expect("worker survived the panic and ran the next job");
|
||
}
|
||
}
|