Fix inconsistent resharding state, SetShardState/AbortTransfer idempotency (#8917)

* [ai] Add integration test for triggering inconsistent resharding state

* [ai] Also add test for resharding down

* Update test

* Resolve resharding idempotency through setting replica states

* Remove resharding down test

* Reformat

* [ai] Remove resharding abort order, abort before setting replica state

* [ai] Resolve test flakiness

* Collapse matches into helper function

* Check preconditions before aborting resharding

* Fix test flakiness by not waiting for a dead node

* Abort resharding before aborting transfer for idempotency

* Update comment

* Release shard holder lock on transfer/reshard abort to prevent deadlock

* Remove now unused shard holder parameter

* Split handle_replica_changes to eliminate need for juggling locking

* In resharding tests, import all utils to enable every_test cleanup

Resolves flakiness I've been seeing in
test_set_replica_dead_clears_resharding_state test
This commit is contained in:
Tim Visée
2026-05-08 16:29:34 +02:00
committed by timvisee
parent 570546c23e
commit ca1971ad4f
8 changed files with 263 additions and 74 deletions

View File

@@ -213,7 +213,7 @@ impl Collection {
/// Handle replica changes
///
/// add and remove replicas from replica set
/// Remove replicas from replica set
pub async fn handle_replica_changes(
&self,
replica_changes: Vec<Change>,
@@ -223,17 +223,17 @@ impl Collection {
}
let shard_holder = self.shards_holder.read().await;
let mut to_remove = Vec::with_capacity(replica_changes.len());
for change in replica_changes {
let (shard_id, peer_id) = match change {
Change::Remove(shard_id, peer_id) => (shard_id, peer_id),
};
let Some(replica_set) = shard_holder.get_shard(shard_id) else {
let Some(replica_set) = shard_holder.get_shard(shard_id).cloned() else {
return Err(CollectionError::bad_request(format!(
"Shard {} of {} not found",
shard_id,
self.name()
"Shard {shard_id} of {} not found",
self.name(),
)));
};
@@ -268,9 +268,15 @@ impl Collection {
.get_transfers(|transfer| transfer.from == peer_id || transfer.to == peer_id)
};
// ...and cancel transfer tasks and remove transfers from internal state
to_remove.push((replica_set, peer_id, transfers));
}
// Must release shard holder lock for abort_shard_transfer_and_resharding
drop(shard_holder);
for (replica_set, peer_id, transfers) in to_remove {
for transfer in transfers {
self.abort_shard_transfer_and_resharding(transfer.key(), Some(&shard_holder))
self.abort_shard_transfer_and_resharding(transfer.key())
.await?;
}
@@ -283,6 +289,7 @@ impl Collection {
// the transfer should be cancelled (see the block right above this comment),
// so no special handling is needed.
}
Ok(())
}

View File

@@ -402,8 +402,8 @@ impl Collection {
new_state: ReplicaState,
from_state: Option<ReplicaState>,
) -> CollectionResult<()> {
let shard_holder = self.shards_holder.read().await;
let replica_set = shard_holder
let mut shard_holder = self.shards_holder.read().await;
let mut replica_set = shard_holder
.get_shard(shard_id)
.ok_or_else(|| shard_not_found_error(shard_id))?;
@@ -451,14 +451,46 @@ impl Collection {
)));
}
// Abort resharding *before* persisting the new replica state.
//
// If we did this after the persist, a crash between the durable
// `replica_state.json` write and `abort_resharding` would leave
// `resharding_state.json` stuck `Some` on the replaying peer (a
// `current_state`-based gate reads `Dead` from disk on replay and
// skips the abort). Doing it first means a partial-apply crash
// either re-runs the whole entry (replica state still reads as a
// resharding state, gate fires, the idempotent abort re-runs) or no
// longer needs the abort to fire at all (resharding_state already
// cleared by the prior attempt). Either way, every peer converges.
//
// Covers both directions: up (`Resharding`) and down (`ReshardingScaleDown`).
// For up, `abort_resharding` drops the shard entirely; we detect the
// missing shard below and skip the persist. For down, the shard
// stays, abort just reverts the receiver back to `Active`, and the
// persist below overwrites that with `Dead`.
if new_state == ReplicaState::Dead && current_state.is_some_and(|s| s.is_resharding()) {
let resharding_state = shard_holder.resharding_state.read().clone();
if let Some(state) = resharding_state {
// `abort_resharding` grabs the shard_holder write lock internally.
drop(shard_holder);
self.abort_resharding(state.key(), false).await?;
// For up direction, abort dropped the shard. If it's gone, the
// replica we'd be deactivating no longer exists on this peer
// and the rest of the function is a no-op.
shard_holder = self.shards_holder.read().await;
let Some(rs) = shard_holder.get_shard(shard_id) else {
return Ok(());
};
replica_set = rs;
}
}
// Update replica status
replica_set
.ensure_replica_with_state(peer_id, new_state)
.await?;
if new_state == ReplicaState::Dead {
let resharding_state = shard_holder.resharding_state.read().clone();
let all_nodes_fixed_cancellation = self
.channel_service
.all_peers_at_version(&ABORT_TRANSFERS_ON_SHARD_DROP_FIX_FROM_VERSION);
@@ -476,36 +508,11 @@ impl Collection {
// Functions below lock `shard_holder`!
drop(shard_holder);
let mut abort_resharding_result = CollectionResult::Ok(());
// Abort resharding, if resharding shard is marked as `Dead`.
//
// This branch should only be triggered, if resharding is currently at `MigratingPoints`
// stage, because target shard should be marked as `Active`, when all resharding transfers
// are successfully completed, and so the check *right above* this one would be triggered.
//
// So, if resharding reached `ReadHashRingCommitted`, this branch *won't* be triggered,
// and resharding *won't* be cancelled. The update request should *fail* with "failed to
// update all replicas of a shard" error.
//
// If resharding reached `ReadHashRingCommitted`, and this branch is triggered *somehow*,
// then `Collection::abort_resharding` call should return an error, so no special handling
// is needed.
let is_resharding = current_state
.as_ref()
.is_some_and(ReplicaState::is_resharding);
if is_resharding && let Some(state) = resharding_state {
abort_resharding_result = self.abort_resharding(state.key(), false).await;
}
// Terminate transfer if source or target replicas are now dead
for transfer in related_transfers {
self.abort_shard_transfer_and_resharding(transfer.key(), None)
self.abort_shard_transfer_and_resharding(transfer.key())
.await?;
}
// Propagate resharding errors now
abort_resharding_result?;
}
// If not initialized yet, we need to check if it was initialized by this call
@@ -668,7 +675,7 @@ impl Collection {
}
for transfer in self.get_related_transfers(peer_id).await {
self.abort_shard_transfer_and_resharding(transfer.key(), None)
self.abort_shard_transfer_and_resharding(transfer.key())
.await?;
}

View File

@@ -436,34 +436,35 @@ impl Collection {
pub async fn abort_shard_transfer_and_resharding(
&self,
transfer_key: ShardTransferKey,
shard_holder: Option<&ShardHolder>,
) -> CollectionResult<()> {
let mut shard_holder_guard = None;
// Look up transfer and any resharding state we need to abort
let resharding_state = {
let shard_holder = self.shards_holder.read().await;
let shard_holder = match shard_holder {
Some(shard_holder) => shard_holder,
None => shard_holder_guard.insert(self.shards_holder.read().await),
let Some(transfer) = shard_holder.get_transfer(&transfer_key) else {
return Ok(());
};
if transfer.is_resharding() {
shard_holder.resharding_state.read().clone()
} else {
None
}
};
// Abort resharding before the transfer to be idempotent
if let Some(state) = resharding_state {
self.abort_resharding(state.key(), false).await?;
}
// Resharding may already have aborted the transfer so we check it again
let shard_holder = self.shards_holder.read().await;
let Some(transfer) = shard_holder.get_transfer(&transfer_key) else {
return Ok(());
};
let is_resharding_transfer = transfer.is_resharding();
self.abort_shard_transfer(transfer, shard_holder).await?;
if is_resharding_transfer {
let resharding_state = shard_holder.resharding_state.read().clone();
// `abort_resharding` locks `shard_holder`!
drop(shard_holder_guard);
if let Some(state) = resharding_state {
self.abort_resharding(state.key(), false).await?;
}
}
Ok(())
self.abort_shard_transfer(transfer, &shard_holder).await
}
/// Initiate local partial shard

View File

@@ -743,9 +743,21 @@ impl ShardReplicaSet {
self.shard_id,
);
// Deactivate replica in consensus if it matches the state we expect
// Always deactivate the replica if its in a shard transfer related state
let from_state = Some(peer_state).filter(|state| !state.is_partial_or_recovery());
// Deactivate replica in consensus if it matches the state we expect.
// We filter out transient transfer states (Partial, Recovery, etc.)
// because those can change rapidly between proposal and apply, and
// we want the deactivation to go through even if the state has
// moved on. Resharding/ReshardingScaleDown are stable through the
// resharding lifecycle (only flipped by finish_migrating_points
// → Active or by this same deactivation flow → Dead), so we keep
// them populated. Carrying the captured pre-state in the consensus
// log entry is what lets the local `abort_resharding` side-effect
// in `Collection::set_shard_replica_state` re-fire correctly on
// replay — without it, replay reads `current_state` as Dead from
// disk and silently skips the side-effect, leaving
// resharding_state stale.
let from_state = Some(peer_state)
.filter(|state| !state.is_partial_or_recovery() || state.is_resharding());
self.add_locally_disabled(Some(state), *peer_id, from_state);
}

View File

@@ -576,7 +576,7 @@ impl TableOfContent {
}
ShardTransferOperations::RecoveryToPartial(transfer)
| ShardTransferOperations::SnapshotRecovered(transfer) => {
// Validate transfer exists to prevent double handling
// Validate transfer exists
transfer::helpers::validate_transfer_exists(
&transfer,
&collection.state().await.transfers,
@@ -635,7 +635,7 @@ impl TableOfContent {
)?;
log::warn!("Aborting shard transfer: {reason}");
collection
.abort_shard_transfer_and_resharding(transfer, None)
.abort_shard_transfer_and_resharding(transfer)
.await?;
}
};

View File

@@ -96,6 +96,7 @@ def upsert_random_points(
num_cities=None,
headers={},
extra_payload=None,
timeout=None,
):
extra_payload = extra_payload or {}
@@ -129,6 +130,7 @@ def upsert_random_points(
"shard_key": shard_key,
},
headers=headers,
timeout=timeout,
)
if fail_on_error:
assert_http_ok(r_batch)

View File

@@ -8,17 +8,7 @@ import requests
from .assertions import assert_http_ok
from .fixtures import create_collection
from .utils import (
get_cluster_info,
get_collection_cluster_info,
processes,
start_cluster,
start_peer,
wait_collection_exists_and_active_on_all_peers,
wait_for,
wait_for_collection_resharding_operations_count,
wait_for_peer_online,
)
from .utils import *
COLLECTION = "test_resharding_replay"

View File

@@ -0,0 +1,170 @@
"""SetShardReplicaState(Resharding -> Dead) must clear resharding_state on every peer.
Reorder fix in `Collection::set_shard_replica_state`: when this entry deactivates
a `Resharding` replica, the local `abort_resharding` runs *before* the durable
replica_state write. With the old order (persist first, abort after), a crash
between the two left `resharding_state` stuck `Some` on the replaying peer; the
replay's gate read `current_state` back as `Dead` from disk and silently
skipped the abort. With the reorder, that intermediate state is unreachable —
either the abort completes (and resharding_state is cleared) before the persist
runs, or the entry hasn't been "applied" from consensus' point of view and
re-runs from the top on replay.
The test exercises the natural failure-detection path: kill a peer holding a
`Resharding` replica, drive updates that fail on it, and verify every peer's
`resharding_operations` clears — including target after it catches up.
"""
import json
import pathlib
import requests
from .assertions import assert_http_ok
from .fixtures import create_collection, upsert_random_points
from .utils import *
COLLECTION = "test_resharding_set_replica_dead_replay"
def _all_replicas(info):
for local in info["local_shards"]:
yield {**local, "peer_id": info["peer_id"]}
for remote in info["remote_shards"]:
yield remote
def test_set_replica_dead_clears_resharding_state(tmp_path: pathlib.Path):
peer_uris, peer_dirs, bootstrap_uri = start_cluster(tmp_path, 3)
target_idx = 2
target_dir = peer_dirs[target_idx]
peer_ids = [get_cluster_info(uri)["peer_id"] for uri in peer_uris]
target_peer_id = peer_ids[target_idx]
# 2 shards (RF=1). Resharding up adds shard 2 on the target peer in
# `Resharding` state. We then drive a resharding-stream-records transfer
# from an existing shard to shard 2 on a non-target peer so shard 2 has
# *two* `Resharding` replicas — without the second replica, the
# failure-detection path's "last active peer" guard refuses to propose
# `SetShardReplicaState(Dead)` and the deactivation never happens.
create_collection(peer_uris[0], collection=COLLECTION, shard_number=2, replication_factor=1)
wait_collection_exists_and_active_on_all_peers(COLLECTION, peer_uris)
new_shard_dir = target_dir / "storage" / "collections" / COLLECTION / "2"
new_replica_state = new_shard_dir / "replica_state.json"
resp = requests.post(
f"{peer_uris[0]}/collections/{COLLECTION}/cluster",
json={"start_resharding": {"direction": "up", "peer_id": target_peer_id, "shard_key": None}},
)
assert_http_ok(resp)
for uri in peer_uris:
wait_for_collection_resharding_operations_count(uri, COLLECTION, 1)
wait_for(lambda: new_replica_state.exists())
wait_for(
lambda: json.loads(new_replica_state.read_text())["peers"].get(str(target_peer_id))
== "Resharding"
)
# Pick a transfer source on a non-target peer (auto-distribution may put
# shard 0 on the target peer). Stream from that peer's existing shard
# into shard 2 on the third peer, giving shard 2 a second Resharding
# replica there.
info = get_collection_cluster_info(peer_uris[0], COLLECTION)
source_idx = source_peer_id = source_shard = None
for r in _all_replicas(info):
replica_idx = peer_ids.index(r["peer_id"])
if replica_idx == target_idx:
continue
source_idx = replica_idx
source_peer_id = r["peer_id"]
source_shard = r["shard_id"]
break
assert source_idx is not None, "no non-target peer holds an existing shard"
extra_idx = next(i for i in range(3) if i != target_idx and i != source_idx)
extra_peer_id = peer_ids[extra_idx]
driver_uri = peer_uris[source_idx]
resp = requests.post(
f"{peer_uris[0]}/collections/{COLLECTION}/cluster",
json={
"replicate_shard": {
"from_peer_id": source_peer_id,
"to_peer_id": extra_peer_id,
"shard_id": source_shard,
"to_shard_id": 2,
"method": "resharding_stream_records",
}
},
)
assert_http_ok(resp)
wait_for_collection_shard_transfers_count(peer_uris[0], COLLECTION, 0)
# Kill target. Updates routed to shard 2 will fail on target's Resharding
# replica; the second peer's Resharding replica keeps the shard "live"
# in active_or_resharding_peers so the failure handler is permitted to
# propose `SetShardReplicaState(Resharding -> Dead)` for target.
p = processes.pop(target_idx)
restart_port = p.p2p_port
p.kill()
# Running peers apply `SetShardReplicaState(Dead)`; the local
# `abort_resharding` runs first (per the reorder), clearing
# `resharding_state` before the replica state write.
#
# With weak ordering the upsert returns as soon as the local replica
# acks, leaving the fan-out failure to target's dead replica to land
# asynchronously — a single batch sometimes returns 200 without ever
# observing the dead target. Drive a fresh batch on each poll until
# every running peer has applied the deactivation and cleared its
# `resharding_state`.
#
# Cap each upsert at 5s so a stuck gRPC channel to the dead peer can't
# block a single iteration past the wait_for budget — without this,
# `requests.put` has no timeout and a hung fan-out leaves wait_for
# unable to iterate at all (test hangs until pytest's outer deadline
# rather than failing within 60s).
running_uris = [uri for idx, uri in enumerate(peer_uris) if idx != target_idx]
def all_running_peers_aborted_resharding() -> bool:
try:
upsert_random_points(
driver_uri,
num=100,
collection_name=COLLECTION,
fail_on_error=False,
ordering="weak",
timeout=5,
)
except requests.RequestException:
pass
return all(
not get_collection_cluster_info(uri, COLLECTION).get("resharding_operations")
for uri in running_uris
)
wait_for(
all_running_peers_aborted_resharding,
wait_for_timeout=60,
wait_for_interval=1,
)
# Restart target. It catches up via raft, applies the same entry, and
# the same reorder makes it clear `resharding_state` and drop shard 2
# on this peer too.
peer_uris[target_idx] = start_peer(
target_dir, f"peer_catchup_{target_idx}.log", bootstrap_uri, port=restart_port
)
catchup_uri = peer_uris[target_idx]
wait_for_peer_online(catchup_uri)
wait_for_collection_resharding_operations_count(catchup_uri, COLLECTION, 0)
# All peers must agree — no divergence in `resharding_operations`.
for uri in peer_uris:
info = get_collection_cluster_info(uri, COLLECTION)
assert not info.get("resharding_operations"), (
f"Resharding state divergence on peer {uri}: "
f"resharding_operations={info.get('resharding_operations')!r}"
)