use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use azure_core::http::headers::{AsHeaders, HeaderName, HeaderValue};
use futures::future::{pending, select, Either, Future};
use crate::{
diagnostics::{DiagnosticsContextBuilder, ExecutionContext, PipelineType, TransportSecurity},
driver::{
routing::{
can_circuit_breaker_trigger_failover, is_eligible_for_ppaf, is_eligible_for_ppcb,
partition_endpoint_state::HealthStatus, partition_key_range_id::PartitionKeyRangeId,
remove_probe_succeeded_entry, session_manager::SessionManager, AccountEndpointState,
CosmosEndpoint, LocationEffect, LocationSnapshot, LocationStateStore,
},
transport::CosmosTransport,
},
models::{
cosmos_headers::QUERY_CONTENT_TYPE, request_header_names, AccountEndpoint, ActivityId,
CosmosOperation, CosmosResponse, Credential, DefaultConsistencyLevel,
EffectivePartitionKey, OperationType, SessionToken, SubStatusCode,
},
options::{
HedgeThreshold, OperationOptionsView, ReadConsistencyStrategy, Region,
ResolvedThroughputControl,
},
};
use super::{
components::{
OperationAction, OperationRetryState, RoutingDecision, TransportMode, TransportOutcome,
TransportRequest, TransportResult, DEFAULT_MAX_THROTTLE_ATTEMPTS,
DEFAULT_MAX_THROTTLE_WAIT,
},
hedging_diagnostics::{HedgeDiagnostics, HedgingStrategyConfig},
hedging_eligibility::evaluate_hedge_eligibility,
retry_evaluation::{
build_service_error, evaluate_hedge_leg_effects, evaluate_transport_result,
is_region_confirming_status, partition_effects_for_deferral,
},
};
use crate::driver::transport::{
transport_pipeline::{execute_transport_pipeline, TransportPipelineContext},
AuthorizationContext,
};
#[derive(Debug, Clone, Default)]
pub(crate) struct OperationOverrides {
pub feed_range: Option<crate::models::FeedRange>,
pub partition_key_range_id: Option<String>,
pub partition_key: Option<crate::models::PartitionKey>,
pub continuation: Option<String>,
}
impl OperationOverrides {
pub fn apply_headers(
&self,
headers: &mut azure_core::http::headers::Headers,
) -> crate::error::Result<()> {
if let Some(feed_range) = &self.feed_range {
if feed_range.min_inclusive() != &EffectivePartitionKey::MIN {
headers.insert(
HeaderName::from_static(request_header_names::START_EPK),
HeaderValue::from(feed_range.min_inclusive().as_str().to_owned()),
);
}
if feed_range.max_exclusive() != &EffectivePartitionKey::MAX {
headers.insert(
HeaderName::from_static(request_header_names::END_EPK),
HeaderValue::from(feed_range.max_exclusive().as_str().to_owned()),
);
}
headers.insert(
HeaderName::from_static(request_header_names::READ_FEED_KEY_TYPE),
HeaderValue::from_static("EffectivePartitionKey"),
);
}
if let Some(pk_range_id) = &self.partition_key_range_id {
headers.insert(
HeaderName::from_static(request_header_names::PARTITION_KEY_RANGE_ID),
HeaderValue::from(pk_range_id.clone()),
);
}
if let Some(pk) = &self.partition_key {
let pk_headers = pk.as_headers()?;
for (name, value) in pk_headers {
headers.insert(name, value);
}
}
if let Some(continuation) = &self.continuation {
headers.insert(
HeaderName::from_static(request_header_names::CONTINUATION),
HeaderValue::from(continuation.clone()),
);
}
Ok(())
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn execute_operation_pipeline(
operation: &CosmosOperation,
overrides: OperationOverrides,
options: &OperationOptionsView<'_>,
custom_headers: Option<&std::collections::HashMap<HeaderName, HeaderValue>>,
location_state_store: &LocationStateStore,
transport: &CosmosTransport,
account_endpoint: &AccountEndpoint,
credential: &Credential,
user_agent: &azure_core::http::headers::HeaderValue,
activity_id: &ActivityId,
pipeline_type: PipelineType,
transport_security: TransportSecurity,
diagnostics: DiagnosticsContextBuilder,
session_manager: &SessionManager,
account_default_consistency: DefaultConsistencyLevel,
throughput_control: Option<ResolvedThroughputControl>,
pre_resolved_pk_range_id: Option<PartitionKeyRangeId>,
) -> crate::error::Result<CosmosResponse> {
let mut diagnostics = diagnostics;
let location_snapshot = location_state_store.snapshot();
let max_failover_retries = options.max_failover_retry_count().copied().unwrap_or(3);
let throttling_retry_options = options.throttling_retry_options();
let max_throttle_attempts = throttling_retry_options
.max_retry_count()
.copied()
.unwrap_or(DEFAULT_MAX_THROTTLE_ATTEMPTS);
let max_throttle_wait_time = throttling_retry_options
.max_retry_wait_time()
.copied()
.unwrap_or(DEFAULT_MAX_THROTTLE_WAIT);
let session_capturing_disabled = options
.session_capturing_disabled()
.copied()
.unwrap_or(false);
let read_consistency_strategy = options
.read_consistency_strategy()
.copied()
.unwrap_or(ReadConsistencyStrategy::Default);
let session_consistency_active = !session_capturing_disabled
&& read_consistency_strategy.is_session_effective(account_default_consistency);
let max_session_retries = options
.max_session_retry_count()
.copied()
.unwrap_or_else(|| {
if location_snapshot.account.multiple_write_locations_enabled {
let endpoints_len = location_snapshot
.account
.preferred_endpoints(operation.is_read_only())
.len();
endpoints_len as u32
} else {
2
}
});
let mut retry_state = OperationRetryState::initial(
location_snapshot.account.generation,
location_snapshot.account.multiple_write_locations_enabled,
options
.excluded_regions()
.map(|r| r.0.clone())
.unwrap_or_default(),
max_failover_retries,
max_session_retries,
);
retry_state.partition_key_range_id = pre_resolved_pk_range_id;
retry_state.ppaf_write_retry_allowed = location_snapshot
.partitions
.per_partition_automatic_failover_enabled
&& !location_snapshot.account.multiple_write_locations_enabled
&& operation
.resource_type()
.is_partitioned(operation.operation_type());
retry_state.ppcb_active = location_snapshot
.partitions
.per_partition_circuit_breaker_enabled
&& location_snapshot.account.preferred_read_endpoints.len() > 1;
retry_state.is_dataplane = pipeline_type.is_data_plane();
let configured_request_timeout = options.end_to_end_latency_policy().map(|p| p.timeout());
let deadline = configured_request_timeout.map(|t| Instant::now() + t);
loop {
let location = location_state_store.snapshot();
let routing = resolve_endpoint(
operation,
&retry_state,
&location,
pipeline_type.is_data_plane(),
location_state_store.endpoint_unavailability_ttl(),
);
tracing::debug!(routing_decision = %routing, "routing decision made");
if retry_state.failover_retry_count == 0 && retry_state.session_token_retry_count == 0 {
if let Some(upgrade) = evaluate_hedge_eligibility(
operation,
options,
&location.account,
&routing,
configured_request_timeout,
) {
let attempt_ctx = AttemptContext {
operation,
overrides: &overrides,
custom_headers,
transport,
account_endpoint,
credential,
user_agent,
activity_id,
pipeline_type,
transport_security,
session_manager,
session_consistency_active,
options,
throughput_control,
deadline,
configured_request_timeout,
can_use_multiple_write_locations: retry_state.can_use_multiple_write_locations,
hub_region_processing_only_initial: retry_state.hub_region_processing_only,
partition_key_range_id: retry_state.partition_key_range_id.clone(),
location_state_store,
};
retry_state.hedge_already_fired = true;
match execute_hedged(
&attempt_ctx,
&routing,
&upgrade.secondary_routing,
upgrade.threshold,
upgrade.strategy_config,
diagnostics,
&retry_state,
)
.await
{
HedgedRaceResult::Terminal(result) => return result,
HedgedRaceResult::BothTransient {
primary_region,
secondary_region,
strategy_config,
primary_region_for_diag,
secondary_region_for_diag,
last_error,
diagnostics: returned_diagnostics,
partition_key_range_id: race_pk_range_id,
observed_session_unavailable: race_observed_1002,
} => {
tracing::warn!(
activity_id = %activity_id,
primary_region = ?primary_region.as_ref().map(crate::options::Region::as_str),
secondary_region = ?secondary_region.as_ref().map(crate::options::Region::as_str),
"hedge race: both transient on STAGE 2b; attempting failover fallback",
);
diagnostics = returned_diagnostics;
if retry_state.partition_key_range_id.is_none() {
retry_state.partition_key_range_id = race_pk_range_id;
}
propagate_hedge_session_unavailable(&mut retry_state, race_observed_1002);
if let Err(e) = try_advance_after_both_transient(
&mut retry_state,
&location,
operation.is_read_only(),
primary_region.as_ref(),
secondary_region.as_ref(),
last_error,
) {
diagnostics.set_hedge_diagnostics(HedgeDiagnostics::both_transient(
strategy_config,
primary_region_for_diag,
secondary_region_for_diag,
false,
));
let diagnostics_ctx = Arc::new(diagnostics.complete());
return Err(crate::error::CosmosErrorBuilder::from_error(e)
.with_diagnostics(diagnostics_ctx)
.build());
}
diagnostics = enforce_deadline_or_timeout(deadline, options, diagnostics)?;
continue;
}
}
}
}
let execution_context = compute_execution_context(&retry_state);
let ctx = TransportRequestContext {
routing: &routing,
activity_id,
execution_context,
deadline,
resolved_session_token: session_consistency_active
.then(|| {
session_manager.resolve_session_token(
operation,
operation.request_headers().session_token.as_ref(),
)
})
.flatten(),
throughput_control,
};
let mut transport_request =
build_transport_request(operation, &overrides, custom_headers, &ctx)?;
apply_hub_region_header(&mut transport_request, &retry_state);
apply_tentative_writes_header(
&mut transport_request,
operation,
location.account.multiple_write_locations_enabled,
);
apply_optional_request_headers(&mut transport_request, operation, options);
tracing::trace!(
method = ?transport_request.method,
url = %transport_request.url,
"transport request created");
let selected_transport = match pipeline_type {
PipelineType::DataPlane => {
transport.get_dataplane_transport(account_endpoint, routing.transport_mode)?
}
PipelineType::Metadata => transport.get_metadata_transport(account_endpoint)?,
};
let result = execute_transport_pipeline(
transport_request,
&TransportPipelineContext {
transport: &selected_transport,
allow_sent_transport_retry: operation.is_read_only() || operation.is_idempotent(),
credential,
user_agent,
pipeline_type,
transport_security,
endpoint_key: routing.endpoint.endpoint_key(),
max_throttle_attempts,
max_throttle_wait_time,
},
&mut diagnostics,
)
.await;
if retry_state.partition_key_range_id.is_none() {
if let Some(headers) = result.cosmos_headers() {
if let Some(pk_range_id) = headers.partition_key_range_id.as_deref() {
retry_state.partition_key_range_id =
Some(PartitionKeyRangeId::from(pk_range_id.to_owned()));
}
}
}
if session_consistency_active {
if let Some(cosmos_headers) = result.cosmos_headers() {
if should_capture_session_token_from_status(
cosmos_headers.substatus.as_ref(),
&result.outcome,
) {
session_manager.capture_session_token(operation, cosmos_headers);
}
}
}
let (action, effects) =
evaluate_transport_result(operation, &routing.endpoint, result, &retry_state);
let action = if retry_state.hedge_already_fired {
action
} else {
maybe_upgrade_to_hedge(
action,
operation,
options,
&location.account,
&routing,
configured_request_timeout,
)
};
let (immediate_effects, deferred_effects) = partition_effects_for_deferral(
operation.is_read_only(),
retry_state.can_use_multiple_write_locations,
retry_state.ppaf_write_retry_allowed,
effects,
);
retry_state.pending_write_effects.extend(deferred_effects);
location_state_store.apply(&immediate_effects).await;
match action {
OperationAction::Complete(result) => {
flush_pending_write_effects(&mut retry_state, location_state_store).await;
try_cleanup_probe_candidate(&retry_state, location_state_store);
return build_cosmos_response(result, diagnostics);
}
OperationAction::FailoverRetry { new_state, delay } => {
tracing::debug!(
activity_id = %activity_id,
failover_attempt = new_state.failover_retry_count,
delay = ?delay,
effects = ?immediate_effects,
deferred_effects = retry_state.pending_write_effects.len(),
"failover retry triggered",
);
apply_failover_delay(delay).await;
advance_to_next_attempt(
&mut retry_state,
new_state,
location_state_store,
operation.is_read_only(),
);
diagnostics = enforce_deadline_or_timeout(deadline, options, diagnostics)?;
}
OperationAction::SessionRetry { new_state } => {
advance_to_next_attempt(
&mut retry_state,
new_state,
location_state_store,
operation.is_read_only(),
);
diagnostics = enforce_deadline_or_timeout(deadline, options, diagnostics)?;
}
OperationAction::Abort { error } => {
let cosmos_status = error.status();
let confirming = is_region_confirming_status(&cosmos_status);
if confirming {
flush_pending_write_effects(&mut retry_state, location_state_store).await;
} else {
retry_state.pending_write_effects.clear();
}
tracing::error!(
activity_id = %activity_id,
status = ?cosmos_status,
error = %error,
operation_type = ?operation.operation_type(),
resource_type = ?operation.resource_type(),
is_read_only = operation.is_read_only(),
is_idempotent = operation.is_idempotent(),
failover_retries = retry_state.failover_retry_count,
session_retries = retry_state.session_token_retry_count,
pk_range_id = ?retry_state.partition_key_range_id,
"operation aborted",
);
diagnostics
.set_operation_status(cosmos_status.status_code(), cosmos_status.sub_status());
let diagnostics_ctx = Arc::new(diagnostics.complete());
return Err(crate::error::CosmosErrorBuilder::from_error(error)
.with_diagnostics(diagnostics_ctx)
.build());
}
OperationAction::Hedge {
secondary_routing: _pre_advance_secondary,
threshold,
strategy_config,
new_state,
} => {
advance_to_next_attempt(
&mut retry_state,
new_state,
location_state_store,
operation.is_read_only(),
);
let location = location_state_store.snapshot();
let primary_routing = resolve_endpoint(
operation,
&retry_state,
&location,
pipeline_type.is_data_plane(),
location_state_store.endpoint_unavailability_ttl(),
);
let secondary_routing = match evaluate_hedge_eligibility(
operation,
options,
&location.account,
&primary_routing,
configured_request_timeout,
) {
Some(upgrade) => upgrade.secondary_routing,
None => {
tracing::debug!(
activity_id = %activity_id,
"STAGE 7 hedge upgrade: no distinct alternate region after \
advance_to_next_attempt; falling back to non-hedged dispatch",
);
continue;
}
};
let attempt_ctx = AttemptContext {
operation,
overrides: &overrides,
custom_headers,
transport,
account_endpoint,
credential,
user_agent,
activity_id,
pipeline_type,
transport_security,
session_manager,
session_consistency_active,
options,
throughput_control,
deadline,
configured_request_timeout,
can_use_multiple_write_locations: retry_state.can_use_multiple_write_locations,
hub_region_processing_only_initial: retry_state.hub_region_processing_only,
partition_key_range_id: retry_state.partition_key_range_id.clone(),
location_state_store,
};
retry_state.hedge_already_fired = true;
match execute_hedged(
&attempt_ctx,
&primary_routing,
&secondary_routing,
threshold,
strategy_config,
diagnostics,
&retry_state,
)
.await
{
HedgedRaceResult::Terminal(result) => return result,
HedgedRaceResult::BothTransient {
primary_region,
secondary_region,
strategy_config,
primary_region_for_diag,
secondary_region_for_diag,
last_error,
diagnostics: returned_diagnostics,
partition_key_range_id: race_pk_range_id,
observed_session_unavailable: race_observed_1002,
} => {
tracing::warn!(
activity_id = %activity_id,
primary_region = ?primary_region.as_ref().map(crate::options::Region::as_str),
secondary_region = ?secondary_region.as_ref().map(crate::options::Region::as_str),
"hedge race: both transient on STAGE 7; attempting failover fallback",
);
diagnostics = returned_diagnostics;
if retry_state.partition_key_range_id.is_none() {
retry_state.partition_key_range_id = race_pk_range_id;
}
propagate_hedge_session_unavailable(&mut retry_state, race_observed_1002);
if let Err(e) = try_advance_after_both_transient(
&mut retry_state,
&location,
operation.is_read_only(),
primary_region.as_ref(),
secondary_region.as_ref(),
last_error,
) {
diagnostics.set_hedge_diagnostics(HedgeDiagnostics::both_transient(
strategy_config,
primary_region_for_diag,
secondary_region_for_diag,
false,
));
let diagnostics_ctx = Arc::new(diagnostics.complete());
return Err(crate::error::CosmosErrorBuilder::from_error(e)
.with_diagnostics(diagnostics_ctx)
.build());
}
diagnostics = enforce_deadline_or_timeout(deadline, options, diagnostics)?;
continue;
}
}
}
}
}
}
async fn flush_pending_write_effects(
retry_state: &mut OperationRetryState,
location_state_store: &LocationStateStore,
) {
if retry_state.pending_write_effects.is_empty() {
return;
}
let pending = std::mem::take(&mut retry_state.pending_write_effects);
let snapshot = location_state_store.snapshot();
let effects: Vec<LocationEffect> = pending
.into_iter()
.filter(|effect| !is_effect_already_applied(effect, &snapshot))
.collect();
if effects.is_empty() {
return;
}
location_state_store.apply(&effects).await;
}
fn is_effect_already_applied(effect: &LocationEffect, snapshot: &LocationSnapshot) -> bool {
match effect {
LocationEffect::MarkEndpointUnavailable { endpoint, .. } => snapshot
.account
.unavailable_endpoints
.contains_key(endpoint.url()),
LocationEffect::MarkPartitionUnavailable(partition) => {
let Some(pk_range_id) = partition.partition_key_range_id.as_ref() else {
return true;
};
let Some(failed_region) = partition.region.as_ref() else {
return false;
};
let partitions = snapshot.partitions.as_ref();
let already_moved = |entry: &crate::driver::routing::partition_endpoint_state::PartitionFailoverEntry| -> bool {
entry
.current_endpoint
.region()
.is_some_and(|r| r != failed_region)
};
partitions
.failover_overrides
.get(pk_range_id.as_str())
.is_some_and(already_moved)
|| partitions
.circuit_breaker_overrides
.get(pk_range_id.as_str())
.is_some_and(already_moved)
}
LocationEffect::RefreshAccountProperties => false,
}
}
fn resolve_endpoint(
operation: &CosmosOperation,
retry_state: &OperationRetryState,
location: &LocationSnapshot,
prefer_gateway20: bool,
endpoint_unavailability_ttl: Duration,
) -> RoutingDecision {
let account = location.account.as_ref();
let read_only = operation.is_read_only();
let in_flight_failed: Vec<&Region> = if !read_only {
retry_state
.pending_write_effects
.iter()
.filter_map(|e| match e {
LocationEffect::MarkPartitionUnavailable(p) => p.region.as_ref(),
LocationEffect::MarkEndpointUnavailable { endpoint, .. } => endpoint.region(),
LocationEffect::RefreshAccountProperties => None,
})
.collect()
} else {
Vec::new()
};
let primary = preferred_endpoints_for_attempt(account, retry_state, read_only);
let selected = try_select_endpoint(
operation,
retry_state,
account,
primary,
&in_flight_failed,
endpoint_unavailability_ttl,
);
let selected = selected.unwrap_or_else(|| {
try_select_endpoint(
operation,
retry_state,
account,
primary,
&[],
endpoint_unavailability_ttl,
)
.unwrap_or_else(|| {
account
.preferred_write_endpoints
.first()
.expect("preferred_write_endpoints is always non-empty")
.clone()
})
});
debug_assert!(
!selected.is_global(),
"pipeline operation resolved to global endpoint; \
this should never happen — only account-topology fetches \
(which bypass this routing path) may use the global endpoint"
);
let use_gateway20 = selected.uses_gateway20(prefer_gateway20);
let transport_mode = if use_gateway20 {
TransportMode::Gateway20
} else {
TransportMode::Gateway
};
if let Some(pk_range_id) = &retry_state.partition_key_range_id {
let partitions = location.partitions.as_ref();
let is_read = operation.is_read_only();
let is_partitioned = operation
.resource_type()
.is_partitioned(operation.operation_type());
let make_partition_routing = |ep: CosmosEndpoint| -> RoutingDecision {
let ep_use_gw20 = ep.uses_gateway20(prefer_gateway20);
RoutingDecision {
selected_url: ep.selected_url(ep_use_gw20).clone(),
transport_mode: if ep_use_gw20 {
TransportMode::Gateway20
} else {
TransportMode::Gateway
},
endpoint: ep,
}
};
let now = Instant::now();
let region_in_flight_failed =
|ep: &CosmosEndpoint| ep.region().is_some_and(|r| in_flight_failed.contains(&r));
let region_not_in_topology = |ep: &CosmosEndpoint| -> bool {
let Some(region) = ep.region() else {
return false; };
!account
.preferred_read_endpoints
.iter()
.chain(account.preferred_write_endpoints.iter())
.any(|known| known.region() == Some(region))
};
let ppcb_should_skip = |ep: &CosmosEndpoint| -> bool {
if region_in_flight_failed(ep) {
return true;
}
if region_not_in_topology(ep) {
return true;
}
let region = ep.region();
let excluded =
region.is_some_and(|r| retry_state.excluded_regions.iter().any(|e| e == r));
if excluded {
return true;
}
!endpoint_is_available(operation, ep, account, now, endpoint_unavailability_ttl)
};
let ppaf_should_skip = region_in_flight_failed;
if is_eligible_for_ppcb(partitions, account, is_read, is_partitioned) {
if let Some(entry) = partitions.circuit_breaker_overrides.get(pk_range_id) {
if entry.health_status == HealthStatus::ProbeCandidate
&& !ppcb_should_skip(&entry.first_failed_endpoint)
{
return make_partition_routing(entry.first_failed_endpoint.clone());
}
if can_circuit_breaker_trigger_failover(entry, is_read, &partitions.config)
&& !ppcb_should_skip(&entry.current_endpoint)
{
return make_partition_routing(entry.current_endpoint.clone());
}
}
} else if is_eligible_for_ppaf(partitions, account, is_read, is_partitioned) {
if let Some(entry) = partitions.failover_overrides.get(pk_range_id) {
if !ppaf_should_skip(&entry.current_endpoint) {
return make_partition_routing(entry.current_endpoint.clone());
}
}
}
}
RoutingDecision {
selected_url: selected.selected_url(use_gateway20).clone(),
endpoint: selected,
transport_mode,
}
}
fn preferred_endpoints_for_attempt<'a>(
account: &'a AccountEndpointState,
retry_state: &OperationRetryState,
read_only: bool,
) -> &'a [CosmosEndpoint] {
if read_only && retry_state.route_reads_to_write_endpoints() {
&account.preferred_write_endpoints
} else if !read_only && retry_state.ppaf_write_retry_allowed {
&account.preferred_read_endpoints
} else {
account.preferred_endpoints(read_only)
}
}
fn try_select_endpoint(
operation: &CosmosOperation,
retry_state: &OperationRetryState,
account: &AccountEndpointState,
endpoints: &[CosmosEndpoint],
skip_regions: &[&Region],
endpoint_unavailability_ttl: Duration,
) -> Option<CosmosEndpoint> {
if endpoints.is_empty() {
return None;
}
let base_index = if retry_state.location.is_current(account.generation) {
retry_state.location.index()
} else {
0
};
let now = Instant::now();
let len = endpoints.len();
let mut first_unavailable = None;
for i in 0..len {
let candidate = &endpoints[(base_index + i) % len];
let candidate_region = candidate.region();
let excluded =
candidate_region.is_some_and(|r| retry_state.excluded_regions.iter().any(|e| e == r));
if excluded {
continue;
}
let in_skip_set = candidate_region.is_some_and(|r| skip_regions.contains(&r));
if in_skip_set {
continue;
}
if endpoint_is_available(
operation,
candidate,
account,
now,
endpoint_unavailability_ttl,
) {
return Some(candidate.clone());
}
if first_unavailable.is_none() {
first_unavailable = Some(candidate.clone());
}
}
first_unavailable
}
fn endpoint_is_available(
operation: &CosmosOperation,
endpoint: &CosmosEndpoint,
account: &AccountEndpointState,
now: Instant,
endpoint_unavailability_ttl: Duration,
) -> bool {
!account
.unavailable_endpoints
.get(endpoint.url())
.is_some_and(|(marked_at, reason)| {
if operation.is_read_only()
&& matches!(
reason,
crate::driver::routing::UnavailableReason::WriteForbidden
)
{
return false;
}
now.saturating_duration_since(*marked_at) < endpoint_unavailability_ttl
})
}
struct TransportRequestContext<'a> {
routing: &'a RoutingDecision,
activity_id: &'a ActivityId,
execution_context: ExecutionContext,
deadline: Option<Instant>,
resolved_session_token: Option<SessionToken>,
throughput_control: Option<ResolvedThroughputControl>,
}
fn build_transport_request(
operation: &CosmosOperation,
overrides: &OperationOverrides,
custom_headers: Option<&std::collections::HashMap<HeaderName, HeaderValue>>,
ctx: &TransportRequestContext<'_>,
) -> crate::error::Result<TransportRequest> {
let paths = operation.compute_resource_paths();
let url = {
let mut base = ctx.routing.selected_url.clone();
let request_path = paths.request_path();
let normalized = if request_path.starts_with('/') {
request_path.to_string()
} else if request_path.is_empty() {
String::new()
} else {
format!("/{}", request_path)
};
base.set_path(&normalized);
base
};
let method = operation.operation_type().http_method();
let resource_type = operation.resource_type();
let auth_context = AuthorizationContext::from_paths(method, resource_type, paths);
let mut headers = azure_core::http::headers::Headers::new();
if let Some(custom) = custom_headers {
for (name, value) in custom {
headers.insert(name.clone(), value.clone());
}
}
operation.request_headers().write_to_headers(&mut headers);
if operation.request_headers().activity_id.is_none() {
headers.insert(
HeaderName::from_static("x-ms-activity-id"),
HeaderValue::from(ctx.activity_id.as_str().to_owned()),
);
}
match operation.operation_type() {
OperationType::Upsert => {
headers.insert(
HeaderName::from_static(request_header_names::IS_UPSERT),
HeaderValue::from_static("true"),
);
}
OperationType::Batch => {
headers.insert(
HeaderName::from_static(request_header_names::IS_BATCH_REQUEST),
HeaderValue::from_static("True"),
);
headers.insert(
HeaderName::from_static(request_header_names::BATCH_ATOMIC),
HeaderValue::from_static("True"),
);
headers.insert(
HeaderName::from_static(request_header_names::BATCH_CONTINUE_ON_ERROR),
HeaderValue::from_static("False"),
);
}
OperationType::Query | OperationType::SqlQuery => {
headers.insert(
HeaderName::from_static(request_header_names::IS_QUERY),
HeaderValue::from_static("True"),
);
headers.insert(
azure_core::http::headers::CONTENT_TYPE,
HeaderValue::from_static(QUERY_CONTENT_TYPE),
);
}
OperationType::QueryPlan => {
headers.insert(
HeaderName::from_static(request_header_names::IS_QUERY),
HeaderValue::from_static("True"),
);
headers.insert(
azure_core::http::headers::CONTENT_TYPE,
HeaderValue::from_static(QUERY_CONTENT_TYPE),
);
headers.insert(
HeaderName::from_static(request_header_names::IS_QUERY_PLAN_REQUEST),
HeaderValue::from_static("True"),
);
}
_ => {}
}
#[cfg(feature = "fault_injection")]
{
if let Some(fault_op) =
crate::fault_injection::FaultOperationType::from_operation_and_resource(
&operation.operation_type(),
&operation.resource_type(),
)
{
crate::driver::transport::cosmos_headers::apply_fault_injection_operation_tag(
&mut headers,
fault_op,
);
}
}
overrides.apply_headers(&mut headers)?;
if let Some(token) = &ctx.resolved_session_token {
headers.insert(
request_header_names::SESSION_TOKEN,
HeaderValue::from(token.as_str().to_owned()),
);
}
if let Some(throughput_control) = ctx.throughput_control {
if let Some(priority) = throughput_control.priority_level {
headers.insert(
request_header_names::PRIORITY_LEVEL,
HeaderValue::from(priority.as_str().to_owned()),
);
}
if let Some(bucket) = throughput_control.throughput_bucket {
headers.insert(
request_header_names::THROUGHPUT_BUCKET,
HeaderValue::from(bucket.to_string()),
);
}
}
Ok(TransportRequest {
method,
endpoint: ctx.routing.endpoint.clone(),
url,
headers,
body: operation.body().map(azure_core::Bytes::copy_from_slice),
auth_context,
execution_context: ctx.execution_context,
deadline: ctx.deadline,
})
}
fn build_cosmos_response(
result: Box<TransportResult>,
mut diagnostics: DiagnosticsContextBuilder,
) -> crate::error::Result<CosmosResponse> {
match result.outcome {
TransportOutcome::Success {
status,
cosmos_headers,
body,
} => {
diagnostics.set_operation_status(status.status_code(), status.sub_status());
let diagnostics_ctx = Arc::new(diagnostics.complete());
Ok(CosmosResponse::new(
body,
cosmos_headers,
status,
diagnostics_ctx,
))
}
_ => {
Err(crate::error::CosmosError::builder()
.with_status(crate::error::CosmosStatus::CLIENT_BUILD_RESPONSE_INVOKED_ON_FAILURE)
.with_message("build_cosmos_response called with non-success result")
.build())
}
}
}
fn should_capture_session_token_from_status(
substatus: Option<&SubStatusCode>,
outcome: &TransportOutcome,
) -> bool {
match outcome {
TransportOutcome::Success { .. } => true,
TransportOutcome::HttpError { status, .. } => {
let code = status.status_code();
if code == azure_core::http::StatusCode::Conflict
|| code == azure_core::http::StatusCode::PreconditionFailed
{
return true;
}
if code == azure_core::http::StatusCode::NotFound {
return substatus != Some(&SubStatusCode::READ_SESSION_NOT_AVAILABLE);
}
false
}
_ => false,
}
}
fn compute_execution_context(retry_state: &OperationRetryState) -> ExecutionContext {
if retry_state.failover_retry_count == 0 && retry_state.session_token_retry_count == 0 {
ExecutionContext::Initial
} else if retry_state.session_token_retry_count > 0 {
ExecutionContext::Retry
} else {
ExecutionContext::RegionFailover
}
}
fn apply_hub_region_header(
transport_request: &mut TransportRequest,
retry_state: &OperationRetryState,
) {
if should_emit_hub_region_header(
retry_state.hub_region_processing_only,
retry_state.shared_hub_region_latch.as_ref(),
) {
transport_request.headers.insert(
HeaderName::from_static(request_header_names::HUB_REGION_PROCESSING_ONLY),
HeaderValue::from_static("True"),
);
}
}
fn should_emit_hub_region_header(
per_state_latched: bool,
shared_latch: Option<&Arc<AtomicBool>>,
) -> bool {
per_state_latched || shared_latch.is_some_and(|s| s.load(Ordering::Acquire))
}
fn should_build_shared_hub_region_latch(
pipeline_type: PipelineType,
can_use_multiple_write_locations: bool,
) -> bool {
pipeline_type.is_data_plane() && !can_use_multiple_write_locations
}
fn apply_tentative_writes_header(
transport_request: &mut TransportRequest,
operation: &CosmosOperation,
multiple_write_locations_enabled: bool,
) {
if multiple_write_locations_enabled && !operation.is_read_only() {
transport_request.headers.insert(
HeaderName::from_static(request_header_names::ALLOW_TENTATIVE_WRITES),
HeaderValue::from_static("true"),
);
}
}
fn apply_optional_request_headers(
transport_request: &mut TransportRequest,
operation: &CosmosOperation,
options: &OperationOptionsView<'_>,
) {
if !operation.operation_type().is_read_only()
&& !matches!(
options.content_response_on_write(),
Some(&crate::options::ContentResponseOnWrite::Enabled)
)
{
transport_request.headers.insert(
request_header_names::PREFER,
HeaderValue::from_static("return=minimal"),
);
}
if let Some(custom_headers) = options.custom_headers() {
for (name, value) in custom_headers {
if !transport_request.headers.iter().any(|(n, _)| n == name) {
transport_request
.headers
.insert(name.clone(), value.clone());
}
}
}
}
async fn apply_failover_delay(delay: Option<Duration>) {
let Some(delay) = delay else {
return;
};
if delay.is_zero() {
return;
}
if let Ok(duration) = azure_core::time::Duration::try_from(delay) {
azure_core::sleep(duration).await;
}
}
fn advance_to_next_attempt(
retry_state: &mut OperationRetryState,
new_state: OperationRetryState,
location_state_store: &LocationStateStore,
is_read_only: bool,
) {
let next_location = location_state_store.snapshot();
let endpoints_len =
preferred_endpoints_for_attempt(next_location.account.as_ref(), &new_state, is_read_only)
.len();
let pending = std::mem::take(&mut retry_state.pending_write_effects);
*retry_state = new_state.advance_location(endpoints_len, next_location.account.generation);
retry_state.pending_write_effects = pending;
}
fn enforce_deadline_or_timeout(
deadline: Option<Instant>,
options: &OperationOptionsView<'_>,
mut diagnostics: DiagnosticsContextBuilder,
) -> Result<DiagnosticsContextBuilder, crate::error::CosmosError> {
let Some(d) = deadline else {
return Ok(diagnostics);
};
if Instant::now() < d {
return Ok(diagnostics);
}
let timeout_duration = options
.end_to_end_latency_policy()
.map(|p| p.timeout())
.unwrap_or_default();
diagnostics.set_operation_status(
azure_core::http::StatusCode::RequestTimeout,
Some(SubStatusCode::CLIENT_OPERATION_TIMEOUT),
);
let diagnostics_ctx = Arc::new(diagnostics.complete());
Err(crate::error::CosmosError::builder()
.with_status(crate::models::CosmosStatus::from_parts(
azure_core::http::StatusCode::RequestTimeout,
Some(SubStatusCode::CLIENT_OPERATION_TIMEOUT),
))
.with_message(format!(
"end-to-end operation timeout exceeded ({timeout_duration:?})"
))
.with_diagnostics(diagnostics_ctx)
.build())
}
fn try_cleanup_probe_candidate(
retry_state: &OperationRetryState,
location_state_store: &LocationStateStore,
) {
let Some(pk_range_id) = &retry_state.partition_key_range_id else {
return;
};
let snapshot = location_state_store.snapshot();
let needs_cleanup = snapshot
.partitions
.circuit_breaker_overrides
.get(pk_range_id.as_str())
.is_some_and(|e| e.health_status == HealthStatus::ProbeCandidate);
if !needs_cleanup {
return;
}
location_state_store.apply_partition(|current| {
let is_probe = current
.circuit_breaker_overrides
.get(pk_range_id.as_str())
.is_some_and(|e| e.health_status == HealthStatus::ProbeCandidate);
if is_probe {
remove_probe_succeeded_entry(current, pk_range_id)
} else {
current.clone()
}
});
}
fn record_hedge_outcome(
location_state_store: &LocationStateStore,
outcome: HedgeOutcome,
partition: Option<&PartitionKeyRangeId>,
primary_region: Option<&Region>,
) {
let Some(partition) = partition else {
return;
};
match outcome {
HedgeOutcome::AlternateWin => {
location_state_store.record_consecutive_hedge_win(partition, primary_region)
}
HedgeOutcome::PrimaryWin => {
location_state_store.record_primary_win(partition, primary_region)
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum HedgeOutcome {
PrimaryWin,
AlternateWin,
}
struct AttemptContext<'a> {
operation: &'a CosmosOperation,
overrides: &'a OperationOverrides,
custom_headers: Option<&'a std::collections::HashMap<HeaderName, HeaderValue>>,
transport: &'a CosmosTransport,
account_endpoint: &'a AccountEndpoint,
credential: &'a Credential,
user_agent: &'a azure_core::http::headers::HeaderValue,
activity_id: &'a ActivityId,
pipeline_type: PipelineType,
transport_security: TransportSecurity,
session_manager: &'a SessionManager,
session_consistency_active: bool,
options: &'a OperationOptionsView<'a>,
throughput_control: Option<ResolvedThroughputControl>,
deadline: Option<Instant>,
configured_request_timeout: Option<Duration>,
can_use_multiple_write_locations: bool,
hub_region_processing_only_initial: bool,
partition_key_range_id: Option<PartitionKeyRangeId>,
location_state_store: &'a LocationStateStore,
}
enum HedgeClass {
Final(Box<TransportResult>),
Transient,
}
fn classify_hedge_result(result: crate::error::Result<TransportResult>) -> HedgeClass {
match result {
Ok(tr) => match &tr.outcome {
TransportOutcome::Success { .. } => HedgeClass::Final(Box::new(tr)),
TransportOutcome::HttpError { status, .. } => {
if status.is_final_result() {
HedgeClass::Final(Box::new(tr))
} else {
HedgeClass::Transient
}
}
TransportOutcome::TransportError { .. } | TransportOutcome::DeadlineExceeded { .. } => {
HedgeClass::Transient
}
},
Err(_) => HedgeClass::Transient,
}
}
fn result_is_final(tr: &TransportResult) -> bool {
match &tr.outcome {
TransportOutcome::Success { .. } => true,
TransportOutcome::HttpError { status, .. } => status.is_final_result(),
TransportOutcome::TransportError { .. } | TransportOutcome::DeadlineExceeded { .. } => {
false
}
}
}
fn pk_range_id_from_result(
result: &crate::error::Result<TransportResult>,
) -> Option<PartitionKeyRangeId> {
result
.as_ref()
.ok()
.and_then(|tr| tr.cosmos_headers())
.and_then(|headers| headers.partition_key_range_id.as_deref())
.map(|s| PartitionKeyRangeId::from(s.to_owned()))
}
fn finalize_hedge_attempt(
result: Box<TransportResult>,
diagnostics: DiagnosticsContextBuilder,
) -> crate::error::Result<CosmosResponse> {
match result.outcome {
outcome @ TransportOutcome::Success { .. } => {
build_cosmos_response(Box::new(TransportResult { outcome }), diagnostics)
}
TransportOutcome::HttpError {
status,
cosmos_headers,
body,
..
} => {
tracing::warn!(
activity_id = %diagnostics.activity_id(),
request_count = diagnostics.request_count(),
http_status = u16::from(status.status_code()),
sub_status = ?status.sub_status(),
"cosmos.hedge.terminal_http_error",
);
let diagnostics_ctx = Arc::new(diagnostics.complete());
let base = build_service_error(&status, &cosmos_headers, &body);
Err(crate::error::CosmosErrorBuilder::from_error(base)
.with_diagnostics(diagnostics_ctx)
.build())
}
TransportOutcome::TransportError { error, .. } => {
tracing::warn!(
activity_id = %diagnostics.activity_id(),
request_count = diagnostics.request_count(),
error = %error,
"cosmos.hedge.terminal_transport_error",
);
let diagnostics_ctx = Arc::new(diagnostics.complete());
Err(crate::error::CosmosErrorBuilder::from_error(error)
.with_diagnostics(diagnostics_ctx)
.build())
}
TransportOutcome::DeadlineExceeded { .. } => {
tracing::warn!(
activity_id = %diagnostics.activity_id(),
request_count = diagnostics.request_count(),
"cosmos.hedge.terminal_deadline_exceeded",
);
let mut diagnostics = diagnostics;
diagnostics.set_operation_status(
azure_core::http::StatusCode::RequestTimeout,
Some(SubStatusCode::CLIENT_OPERATION_TIMEOUT),
);
let diagnostics_ctx = Arc::new(diagnostics.complete());
Err(crate::error::CosmosError::builder()
.with_status(crate::models::CosmosStatus::from_parts(
azure_core::http::StatusCode::RequestTimeout,
Some(SubStatusCode::CLIENT_OPERATION_TIMEOUT),
))
.with_message("deadline exceeded during hedged attempt")
.with_diagnostics(diagnostics_ctx)
.build())
}
}
}
fn maybe_upgrade_to_hedge(
action: OperationAction,
operation: &CosmosOperation,
options: &OperationOptionsView<'_>,
account_state: &AccountEndpointState,
primary: &RoutingDecision,
request_timeout: Option<Duration>,
) -> OperationAction {
let new_state = match &action {
OperationAction::FailoverRetry { new_state, .. } => new_state.clone(),
OperationAction::SessionRetry { new_state } => new_state.clone(),
_ => return action,
};
match evaluate_hedge_eligibility(operation, options, account_state, primary, request_timeout) {
Some(upgrade) => {
if new_state.failover_retry_count.saturating_add(2) > new_state.max_failover_retries {
tracing::debug!(
failover_retry_count = new_state.failover_retry_count,
max_failover_retries = new_state.max_failover_retries,
"cosmos.hedge.budget_exhausted_skipping_upgrade",
);
return action;
}
tracing::debug!(
threshold_ms = upgrade.threshold.get().as_millis() as u64,
primary_region = ?primary.endpoint.region().map(crate::options::Region::as_str),
secondary_region = ?upgrade.secondary_routing.endpoint.region().map(crate::options::Region::as_str),
hub_region_processing_only = new_state.hub_region_processing_only,
"cosmos.hedge.enabled_for_operation",
);
OperationAction::Hedge {
secondary_routing: upgrade.secondary_routing,
threshold: upgrade.threshold,
strategy_config: upgrade.strategy_config,
new_state,
}
}
None => action,
}
}
async fn perform_single_attempt(
ctx: &AttemptContext<'_>,
routing: &RoutingDecision,
execution_context: ExecutionContext,
shared_hub_region_latch: Option<&Arc<AtomicBool>>,
diagnostics: &mut DiagnosticsContextBuilder,
) -> crate::error::Result<TransportResult> {
let resolved_session_token = ctx
.session_consistency_active
.then(|| {
ctx.session_manager.resolve_session_token(
ctx.operation,
ctx.operation.request_headers().session_token.as_ref(),
)
})
.flatten();
let request_ctx = TransportRequestContext {
routing,
activity_id: ctx.activity_id,
execution_context,
deadline: ctx.deadline,
resolved_session_token,
throughput_control: ctx.throughput_control,
};
let mut transport_request = build_transport_request(
ctx.operation,
ctx.overrides,
ctx.custom_headers,
&request_ctx,
)?;
if should_emit_hub_region_header(false, shared_hub_region_latch) {
transport_request.headers.insert(
HeaderName::from_static(request_header_names::HUB_REGION_PROCESSING_ONLY),
HeaderValue::from_static("True"),
);
}
apply_optional_request_headers(&mut transport_request, ctx.operation, ctx.options);
let selected_transport = match ctx.pipeline_type {
PipelineType::DataPlane => ctx
.transport
.get_dataplane_transport(ctx.account_endpoint, routing.transport_mode)?,
PipelineType::Metadata => ctx.transport.get_metadata_transport(ctx.account_endpoint)?,
};
let throttling_retry_options = ctx.options.throttling_retry_options();
let max_throttle_attempts = throttling_retry_options
.max_retry_count()
.copied()
.unwrap_or(DEFAULT_MAX_THROTTLE_ATTEMPTS);
let max_throttle_wait_time = throttling_retry_options
.max_retry_wait_time()
.copied()
.unwrap_or(DEFAULT_MAX_THROTTLE_WAIT);
let result = execute_transport_pipeline(
transport_request,
&TransportPipelineContext {
transport: &selected_transport,
allow_sent_transport_retry: ctx.operation.is_read_only()
|| ctx.operation.is_idempotent(),
credential: ctx.credential,
user_agent: ctx.user_agent,
pipeline_type: ctx.pipeline_type,
transport_security: ctx.transport_security,
endpoint_key: routing.endpoint.endpoint_key(),
max_throttle_attempts,
max_throttle_wait_time,
},
diagnostics,
)
.await;
Ok(result)
}
fn capture_session_token_for_winner(ctx: &AttemptContext<'_>, result: &TransportResult) {
if !ctx.session_consistency_active {
return;
}
if let Some(cosmos_headers) = result.cosmos_headers() {
if should_capture_session_token_from_status(
cosmos_headers.substatus.as_ref(),
&result.outcome,
) {
ctx.session_manager
.capture_session_token(ctx.operation, cosmos_headers);
}
}
}
const HARVEST_WINDOW: Duration = Duration::from_millis(50);
const HARVEST_WINDOW_FLOOR: Duration = Duration::from_millis(1);
fn effective_harvest_window(configured_request_timeout: Option<Duration>) -> Duration {
match configured_request_timeout {
Some(t) => HARVEST_WINDOW.min(t / 10).max(HARVEST_WINDOW_FLOOR),
None => HARVEST_WINDOW,
}
}
enum TimerEvent {
ThresholdElapsed,
DeadlineFired,
}
fn deadline_signal(deadline: Option<Instant>) -> Pin<Box<dyn Future<Output = ()> + Send>> {
let Some(d) = deadline else {
return Box::pin(pending::<()>());
};
let remaining_std = d.saturating_duration_since(Instant::now());
match azure_core::time::Duration::try_from(remaining_std) {
Ok(remaining) => Box::pin(azure_core::sleep(remaining)),
Err(_) => Box::pin(futures::future::ready(())),
}
}
async fn harvest_remaining_attempt<F>(
attempt: F,
parent: &mut DiagnosticsContextBuilder,
harvest_window: Duration,
) where
F: Future<
Output = (
crate::error::Result<TransportResult>,
DiagnosticsContextBuilder,
),
> + Unpin
+ Send,
{
let window = match azure_core::time::Duration::try_from(harvest_window) {
Ok(d) => d,
Err(_) => return,
};
let timer = Box::pin(azure_core::sleep(window));
if let Either::Left(((_result, diag), _timer)) = select(attempt, timer).await {
parent.merge_hedge_attempt(diag);
}
}
fn application_cancelled_error(
mut diagnostics: DiagnosticsContextBuilder,
) -> crate::error::CosmosError {
diagnostics.set_operation_status(
azure_core::http::StatusCode::RequestTimeout,
Some(SubStatusCode::CLIENT_OPERATION_TIMEOUT),
);
let diagnostics_ctx = Arc::new(diagnostics.complete());
crate::error::CosmosError::builder()
.with_status(crate::models::CosmosStatus::from_parts(
azure_core::http::StatusCode::RequestTimeout,
Some(SubStatusCode::CLIENT_OPERATION_TIMEOUT),
))
.with_message("operation cancelled by application deadline during cross-region hedging")
.with_diagnostics(diagnostics_ctx)
.build()
}
async fn await_attempt_or_deadline_harvest<F>(
attempt: F,
deadline: Option<Instant>,
parent: &mut DiagnosticsContextBuilder,
harvest_window: Duration,
) -> Option<(
crate::error::Result<TransportResult>,
DiagnosticsContextBuilder,
)>
where
F: Future<
Output = (
crate::error::Result<TransportResult>,
DiagnosticsContextBuilder,
),
> + Unpin
+ Send,
{
let deadline_fut = deadline_signal(deadline);
match select(attempt, deadline_fut).await {
Either::Left((result, _deadline)) => Some(result),
Either::Right(((), remaining)) => {
harvest_remaining_attempt(remaining, parent, harvest_window).await;
None
}
}
}
fn deadline_elapsed(deadline: Option<Instant>) -> bool {
deadline.is_some_and(|d| Instant::now() >= d)
}
#[derive(Debug)]
#[allow(
clippy::large_enum_variant,
reason = "single return value from an async function on a hot path; \
boxing either variant would add a heap allocation per race \
without enabling reuse — at most one `HedgedRaceResult` \
exists at a time per operation."
)]
pub(crate) enum HedgedRaceResult {
Terminal(crate::error::Result<CosmosResponse>),
BothTransient {
primary_region: Option<Region>,
secondary_region: Option<Region>,
strategy_config: HedgingStrategyConfig,
primary_region_for_diag: Region,
secondary_region_for_diag: Region,
last_error: crate::error::CosmosError,
diagnostics: DiagnosticsContextBuilder,
partition_key_range_id: Option<PartitionKeyRangeId>,
observed_session_unavailable: bool,
},
}
async fn apply_hedge_leg_effects(
ctx: &AttemptContext<'_>,
retry_state_snapshot: &OperationRetryState,
endpoint: &CosmosEndpoint,
result: &crate::error::Result<TransportResult>,
shared_hub_region_latch: Option<&Arc<AtomicBool>>,
race_observed_session_unavailable: &mut bool,
) {
let Ok(transport_result) = result.as_ref() else {
return;
};
let eval = evaluate_hedge_leg_effects(
ctx.operation,
endpoint,
retry_state_snapshot,
transport_result,
);
if !eval.effects.is_empty() {
ctx.location_state_store.apply(&eval.effects).await;
}
if eval.observed_session_unavailable {
*race_observed_session_unavailable = true;
if let Some(latch) = shared_hub_region_latch {
latch.store(true, Ordering::Release);
}
}
}
async fn execute_hedged(
ctx: &AttemptContext<'_>,
primary_routing: &RoutingDecision,
secondary_routing: &RoutingDecision,
threshold: HedgeThreshold,
strategy_config: HedgingStrategyConfig,
mut parent_diagnostics: DiagnosticsContextBuilder,
retry_state_snapshot: &OperationRetryState,
) -> HedgedRaceResult {
let primary_region = primary_routing.endpoint.region().cloned();
let secondary_region = secondary_routing.endpoint.region().cloned();
let primary_region_for_diag = primary_region
.clone()
.unwrap_or_else(|| Region::new(HedgeDiagnostics::UNKNOWN_REGION_SENTINEL));
let secondary_region_for_diag = secondary_region
.clone()
.unwrap_or_else(|| Region::new(HedgeDiagnostics::UNKNOWN_REGION_SENTINEL));
tracing::debug!(
activity_id = %ctx.activity_id,
threshold_ms = ?threshold.get().as_millis(),
primary_region = ?primary_region.as_ref().map(|r| r.as_str()),
secondary_region = ?secondary_region.as_ref().map(|r| r.as_str()),
"execute_hedged: launching primary attempt",
);
let primary_diag = parent_diagnostics.clone_for_hedge_attempt();
let primary_attempt = Box::pin(async move {
let mut diag = primary_diag;
let result = perform_single_attempt(
ctx,
primary_routing,
ExecutionContext::Initial,
None,
&mut diag,
)
.await;
(result, diag)
});
let threshold_duration = match azure_core::time::Duration::try_from(threshold.get()) {
Ok(d) => d,
Err(_) => {
parent_diagnostics.set_operation_status(
azure_core::http::StatusCode::InternalServerError,
Some(SubStatusCode::TRANSPORT_GENERATED_503),
);
let diagnostics_ctx = Arc::new(parent_diagnostics.complete());
return HedgedRaceResult::Terminal(Err(crate::error::CosmosError::builder()
.with_status(crate::models::CosmosStatus::TRANSPORT_GENERATED_503)
.with_message("hedge threshold exceeds azure_core::time::Duration range")
.with_diagnostics(diagnostics_ctx)
.build()));
}
};
let threshold_timer = Box::pin(azure_core::sleep(threshold_duration));
let deadline_timer = deadline_signal(ctx.deadline);
let timer_event = Box::pin(async move {
match select(threshold_timer, deadline_timer).await {
Either::Left(((), _)) => TimerEvent::ThresholdElapsed,
Either::Right(((), _)) => TimerEvent::DeadlineFired,
}
});
type PrimaryAttemptFuture<'fut> = Pin<
Box<
dyn Future<
Output = (
crate::error::Result<TransportResult>,
DiagnosticsContextBuilder,
),
> + Send
+ 'fut,
>,
>;
let primary_attempt: PrimaryAttemptFuture<'_> = match select(primary_attempt, timer_event).await
{
Either::Left(((result, diag), _timer)) => {
let primary_was_final = matches!(
&result,
Ok(tr) if result_is_final(tr),
);
if primary_was_final {
parent_diagnostics.merge_hedge_attempt(diag);
parent_diagnostics.set_hedge_diagnostics(HedgeDiagnostics::primary_only(
strategy_config,
primary_region_for_diag.clone(),
));
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: primary won pre-threshold (zero-overhead happy path)",
);
tracing::debug!(
activity_id = %ctx.activity_id,
winner_region = ?primary_region.as_ref().map(crate::options::Region::as_str),
was_hedge = false,
"cosmos.hedge.winner_selected",
);
let pk_range_id_for_feedback =
pk_range_id_from_result(&result).or_else(|| ctx.partition_key_range_id.clone());
record_hedge_outcome(
ctx.location_state_store,
HedgeOutcome::PrimaryWin,
pk_range_id_for_feedback.as_ref(),
primary_region.as_ref(),
);
let tr = result.expect("Ok by primary_was_final guard");
capture_session_token_for_winner(ctx, &tr);
return HedgedRaceResult::Terminal(finalize_hedge_attempt(
Box::new(tr),
parent_diagnostics,
));
}
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: primary completed pre-threshold transient; launching secondary",
);
Box::pin(futures::future::ready((result, diag)))
}
Either::Right((TimerEvent::ThresholdElapsed, remaining_primary)) => remaining_primary,
Either::Right((TimerEvent::DeadlineFired, remaining_primary)) => {
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: deadline fired pre-threshold; harvesting primary",
);
harvest_remaining_attempt(
remaining_primary,
&mut parent_diagnostics,
effective_harvest_window(ctx.configured_request_timeout),
)
.await;
parent_diagnostics.set_hedge_diagnostics(
HedgeDiagnostics::primary_only_deadline_exceeded(
strategy_config,
primary_region_for_diag.clone(),
),
);
return HedgedRaceResult::Terminal(Err(application_cancelled_error(
parent_diagnostics,
)));
}
};
let shared_hub_region_latch = should_build_shared_hub_region_latch(
ctx.pipeline_type,
ctx.can_use_multiple_write_locations,
)
.then(|| Arc::new(AtomicBool::new(ctx.hub_region_processing_only_initial)));
let secondary_shared_latch = shared_hub_region_latch.clone();
let secondary_diag = parent_diagnostics.clone_for_hedge_attempt();
let secondary_attempt = Box::pin(async move {
let mut diag = secondary_diag;
let result = perform_single_attempt(
ctx,
secondary_routing,
ExecutionContext::Hedging,
secondary_shared_latch.as_ref(),
&mut diag,
)
.await;
(result, diag)
});
tracing::debug!(
activity_id = %ctx.activity_id,
shared_hub_region_latch = shared_hub_region_latch.is_some(),
"execute_hedged: threshold elapsed; secondary launched",
);
tracing::debug!(
activity_id = %ctx.activity_id,
threshold_ms = threshold.get().as_millis() as u64,
secondary_region = ?secondary_region.as_ref().map(crate::options::Region::as_str),
"cosmos.hedge.alternate_spawned",
);
let mut captured_pk_range_id: Option<PartitionKeyRangeId> = None;
let mut race_observed_session_unavailable = false;
match select(primary_attempt, secondary_attempt).await {
Either::Left(((primary_result, primary_diag), secondary_remaining)) => {
parent_diagnostics.merge_hedge_attempt(primary_diag);
captured_pk_range_id =
captured_pk_range_id.or_else(|| pk_range_id_from_result(&primary_result));
apply_hedge_leg_effects(
ctx,
retry_state_snapshot,
&primary_routing.endpoint,
&primary_result,
shared_hub_region_latch.as_ref(),
&mut race_observed_session_unavailable,
)
.await;
match classify_hedge_result(primary_result) {
HedgeClass::Final(tr) => {
parent_diagnostics.set_hedge_diagnostics(
HedgeDiagnostics::primary_won_after_hedge(
strategy_config,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
),
);
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: primary won after threshold",
);
tracing::debug!(
activity_id = %ctx.activity_id,
which = "secondary",
target_region = ?secondary_region.as_ref().map(crate::options::Region::as_str),
reason = "primary_won_post_threshold",
"cosmos.hedge.canceled",
);
tracing::debug!(
activity_id = %ctx.activity_id,
winner_region = ?primary_region.as_ref().map(crate::options::Region::as_str),
was_hedge = true,
"cosmos.hedge.winner_selected",
);
record_hedge_outcome(
ctx.location_state_store,
HedgeOutcome::PrimaryWin,
captured_pk_range_id
.as_ref()
.or(ctx.partition_key_range_id.as_ref()),
primary_region.as_ref(),
);
capture_session_token_for_winner(ctx, &tr);
HedgedRaceResult::Terminal(finalize_hedge_attempt(tr, parent_diagnostics))
}
HedgeClass::Transient => {
let Some((secondary_result, secondary_diag)) =
await_attempt_or_deadline_harvest(
secondary_remaining,
ctx.deadline,
&mut parent_diagnostics,
effective_harvest_window(ctx.configured_request_timeout),
)
.await
else {
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: deadline fired awaiting secondary after primary transient",
);
parent_diagnostics.set_hedge_diagnostics(
HedgeDiagnostics::cancelled_awaiting_partner(
strategy_config,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
),
);
return HedgedRaceResult::Terminal(Err(application_cancelled_error(
parent_diagnostics,
)));
};
parent_diagnostics.merge_hedge_attempt(secondary_diag);
captured_pk_range_id =
captured_pk_range_id.or_else(|| pk_range_id_from_result(&secondary_result));
apply_hedge_leg_effects(
ctx,
retry_state_snapshot,
&secondary_routing.endpoint,
&secondary_result,
shared_hub_region_latch.as_ref(),
&mut race_observed_session_unavailable,
)
.await;
match classify_hedge_result(secondary_result) {
HedgeClass::Final(tr) => {
parent_diagnostics.set_hedge_diagnostics(HedgeDiagnostics::hedge_won(
strategy_config,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
));
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: secondary won after primary transient",
);
tracing::debug!(
activity_id = %ctx.activity_id,
winner_region = ?secondary_region.as_ref().map(crate::options::Region::as_str),
was_hedge = true,
"cosmos.hedge.winner_selected",
);
record_hedge_outcome(
ctx.location_state_store,
HedgeOutcome::AlternateWin,
captured_pk_range_id
.as_ref()
.or(ctx.partition_key_range_id.as_ref()),
primary_region.as_ref(),
);
capture_session_token_for_winner(ctx, &tr);
HedgedRaceResult::Terminal(finalize_hedge_attempt(
tr,
parent_diagnostics,
))
}
HedgeClass::Transient => {
finalize_both_transient(
ctx.activity_id,
ctx.deadline,
strategy_config,
primary_region,
secondary_region,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
parent_diagnostics,
captured_pk_range_id,
race_observed_session_unavailable,
)
}
}
}
}
}
Either::Right(((secondary_result, secondary_diag), primary_remaining)) => {
parent_diagnostics.merge_hedge_attempt(secondary_diag);
captured_pk_range_id =
captured_pk_range_id.or_else(|| pk_range_id_from_result(&secondary_result));
apply_hedge_leg_effects(
ctx,
retry_state_snapshot,
&secondary_routing.endpoint,
&secondary_result,
shared_hub_region_latch.as_ref(),
&mut race_observed_session_unavailable,
)
.await;
match classify_hedge_result(secondary_result) {
HedgeClass::Final(tr) => {
parent_diagnostics.set_hedge_diagnostics(HedgeDiagnostics::hedge_won(
strategy_config,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
));
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: secondary won race",
);
tracing::debug!(
activity_id = %ctx.activity_id,
which = "primary",
target_region = ?primary_region.as_ref().map(crate::options::Region::as_str),
reason = "secondary_won_race",
"cosmos.hedge.canceled",
);
tracing::debug!(
activity_id = %ctx.activity_id,
winner_region = ?secondary_region.as_ref().map(crate::options::Region::as_str),
was_hedge = true,
"cosmos.hedge.winner_selected",
);
record_hedge_outcome(
ctx.location_state_store,
HedgeOutcome::AlternateWin,
captured_pk_range_id
.as_ref()
.or(ctx.partition_key_range_id.as_ref()),
primary_region.as_ref(),
);
capture_session_token_for_winner(ctx, &tr);
HedgedRaceResult::Terminal(finalize_hedge_attempt(tr, parent_diagnostics))
}
HedgeClass::Transient => {
let Some((primary_result, primary_diag)) = await_attempt_or_deadline_harvest(
primary_remaining,
ctx.deadline,
&mut parent_diagnostics,
effective_harvest_window(ctx.configured_request_timeout),
)
.await
else {
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: deadline fired awaiting primary after secondary transient",
);
parent_diagnostics.set_hedge_diagnostics(
HedgeDiagnostics::cancelled_awaiting_partner(
strategy_config,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
),
);
return HedgedRaceResult::Terminal(Err(application_cancelled_error(
parent_diagnostics,
)));
};
parent_diagnostics.merge_hedge_attempt(primary_diag);
captured_pk_range_id =
captured_pk_range_id.or_else(|| pk_range_id_from_result(&primary_result));
apply_hedge_leg_effects(
ctx,
retry_state_snapshot,
&primary_routing.endpoint,
&primary_result,
shared_hub_region_latch.as_ref(),
&mut race_observed_session_unavailable,
)
.await;
match classify_hedge_result(primary_result) {
HedgeClass::Final(tr) => {
parent_diagnostics.set_hedge_diagnostics(
HedgeDiagnostics::primary_won_after_hedge(
strategy_config,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
),
);
tracing::debug!(
activity_id = %ctx.activity_id,
"execute_hedged: primary won after secondary transient",
);
tracing::debug!(
activity_id = %ctx.activity_id,
winner_region = ?primary_region.as_ref().map(crate::options::Region::as_str),
was_hedge = true,
"cosmos.hedge.winner_selected",
);
record_hedge_outcome(
ctx.location_state_store,
HedgeOutcome::PrimaryWin,
captured_pk_range_id
.as_ref()
.or(ctx.partition_key_range_id.as_ref()),
primary_region.as_ref(),
);
capture_session_token_for_winner(ctx, &tr);
HedgedRaceResult::Terminal(finalize_hedge_attempt(
tr,
parent_diagnostics,
))
}
HedgeClass::Transient => {
finalize_both_transient(
ctx.activity_id,
ctx.deadline,
strategy_config,
primary_region,
secondary_region,
primary_region_for_diag.clone(),
secondary_region_for_diag.clone(),
parent_diagnostics,
captured_pk_range_id,
race_observed_session_unavailable,
)
}
}
}
}
}
}
}
fn transient_outcome_error(
primary_region: Option<&Region>,
secondary_region: Option<&Region>,
) -> crate::error::CosmosError {
let p = primary_region
.map(Region::as_str)
.unwrap_or(HedgeDiagnostics::UNKNOWN_REGION_SENTINEL);
let s = secondary_region
.map(Region::as_str)
.unwrap_or(HedgeDiagnostics::UNKNOWN_REGION_SENTINEL);
crate::error::CosmosError::builder()
.with_status(crate::models::CosmosStatus::TRANSPORT_GENERATED_503)
.with_message(format!(
"hedging completed without producing a final response \
(primary={p}, secondary={s})"
))
.build()
}
#[allow(
clippy::too_many_arguments,
reason = "All ten values come from `execute_hedged`'s state and \
are passed through to a single helper call site. Bundling \
them into a struct would add an indirection without \
improving readability — the two-region pairing (raw / \
for-diag) is intentional and explicit at the call site."
)]
fn finalize_both_transient(
activity_id: &ActivityId,
deadline: Option<Instant>,
strategy_config: HedgingStrategyConfig,
primary_region: Option<Region>,
secondary_region: Option<Region>,
primary_region_for_diag: Region,
secondary_region_for_diag: Region,
mut parent_diagnostics: DiagnosticsContextBuilder,
partition_key_range_id: Option<PartitionKeyRangeId>,
observed_session_unavailable: bool,
) -> HedgedRaceResult {
let deadline_was_elapsed = deadline_elapsed(deadline);
tracing::warn!(
activity_id = %activity_id,
deadline_elapsed = deadline_was_elapsed,
primary_region = ?primary_region.as_ref().map(Region::as_str),
secondary_region = ?secondary_region.as_ref().map(Region::as_str),
"cosmos.hedge.both_transient",
);
if deadline_was_elapsed {
parent_diagnostics.set_hedge_diagnostics(HedgeDiagnostics::both_transient(
strategy_config,
primary_region_for_diag,
secondary_region_for_diag,
deadline_was_elapsed,
));
tracing::debug!(
activity_id = %activity_id,
"execute_hedged: both transient under elapsed deadline; surfacing app-cancel",
);
HedgedRaceResult::Terminal(Err(application_cancelled_error(parent_diagnostics)))
} else {
tracing::debug!(
activity_id = %activity_id,
primary_region = ?primary_region.as_ref().map(Region::as_str),
secondary_region = ?secondary_region.as_ref().map(Region::as_str),
observed_session_unavailable,
"execute_hedged: both legs transient; bubbling up for failover-loop fallback",
);
HedgedRaceResult::BothTransient {
last_error: transient_outcome_error(primary_region.as_ref(), secondary_region.as_ref()),
primary_region,
secondary_region,
strategy_config,
primary_region_for_diag,
secondary_region_for_diag,
diagnostics: parent_diagnostics,
partition_key_range_id,
observed_session_unavailable,
}
}
}
fn try_advance_after_both_transient(
retry_state: &mut OperationRetryState,
location: &LocationSnapshot,
is_read_only: bool,
primary_region: Option<&Region>,
secondary_region: Option<&Region>,
last_error: crate::error::CosmosError,
) -> Result<(), crate::error::CosmosError> {
let consumed: u32 = 2;
let next_count = retry_state.failover_retry_count.saturating_add(consumed);
if next_count > retry_state.max_failover_retries {
tracing::debug!(
failover_retry_count = retry_state.failover_retry_count,
max_failover_retries = retry_state.max_failover_retries,
"hedge both-transient: failover budget exhausted; surfacing terminal error",
);
return Err(last_error);
}
retry_state.failover_retry_count = next_count;
let endpoints =
preferred_endpoints_for_attempt(location.account.as_ref(), retry_state, is_read_only);
let endpoints_len = endpoints.len();
if endpoints_len > 0 {
let raced_region = |candidate: Option<&Region>| -> bool {
match candidate {
Some(c) => {
primary_region.is_some_and(|p| p == c)
|| secondary_region.is_some_and(|s| s == c)
}
None => false,
}
};
let mut advanced = retry_state
.location
.next_for_generation(endpoints_len, location.account.generation);
for _ in 0..endpoints_len.saturating_sub(1) {
if !raced_region(endpoints[advanced.index()].region()) {
break;
}
advanced = advanced.next_for_generation(endpoints_len, location.account.generation);
}
retry_state.location = advanced;
}
tracing::debug!(
failover_retry_count = retry_state.failover_retry_count,
max_failover_retries = retry_state.max_failover_retries,
"hedge both-transient: failover loop will continue against remaining regions",
);
Ok(())
}
fn propagate_hedge_session_unavailable(
retry_state: &mut OperationRetryState,
observed_session_unavailable: bool,
) {
if !observed_session_unavailable {
return;
}
if retry_state.can_retry_session() {
retry_state.session_token_retry_count =
retry_state.session_token_retry_count.saturating_add(1);
}
if retry_state.hub_region_processing_only {
return;
}
retry_state.hub_region_processing_only = true;
if let Some(shared) = retry_state.shared_hub_region_latch.as_ref() {
shared.store(true, Ordering::Release);
}
tracing::debug!(
session_token_retry_count = retry_state.session_token_retry_count,
max_session_retries = retry_state.max_session_retries,
"hedge both-transient: 1002 observed by at least one leg; \
flipped hub_region_processing_only latch and advanced \
session-retry counter for next attempt",
);
}
#[cfg(test)]
mod tests {
use std::{sync::Arc, time::Duration};
use azure_core::http::headers::HeaderName;
use url::Url;
use super::build_transport_request;
use super::OperationOverrides;
use super::TransportRequestContext;
use crate::{
diagnostics::ExecutionContext,
driver::{
pipeline::components::{RoutingDecision, TransportMode},
routing::{
AccountEndpointState, CosmosEndpoint, LocationEffect, LocationIndex,
LocationSnapshot,
},
},
models::{
request_header_names, AccountReference, ActivityId, ContainerProperties,
ContainerReference, CosmosOperation, DatabaseReference, EffectivePartitionKey,
FeedRange, ItemReference, PartitionKey, PartitionKeyDefinition, SystemProperties,
},
options::{PriorityLevel, ResolvedThroughputControl},
};
fn test_account() -> AccountReference {
AccountReference::with_master_key(
Url::parse("https://test.documents.azure.com:443/").unwrap(),
"test-key",
)
}
fn test_partition_key_definition(path: &str) -> PartitionKeyDefinition {
serde_json::from_str(&format!(r#"{{"paths":["{path}"]}}"#)).unwrap()
}
fn test_container_props() -> ContainerProperties {
ContainerProperties {
id: "testcontainer".into(),
partition_key: test_partition_key_definition("/pk"),
system_properties: SystemProperties::default(),
}
}
fn test_container() -> ContainerReference {
ContainerReference::new(
test_account(),
"testdb",
"testdb_rid",
"testcontainer",
"testcontainer_rid",
&test_container_props(),
)
}
fn test_routing() -> RoutingDecision {
let endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
RoutingDecision {
selected_url: endpoint.url().clone(),
endpoint,
transport_mode: TransportMode::Gateway,
}
}
#[test]
fn apply_headers_pk_range_only_omits_read_key_type() {
let overrides = OperationOverrides {
partition_key_range_id: Some("0".to_string()),
..Default::default()
};
let mut headers = azure_core::http::headers::Headers::new();
overrides
.apply_headers(&mut headers)
.expect("apply_headers should succeed");
assert_eq!(
headers
.get_optional_str(&HeaderName::from_static(
request_header_names::PARTITION_KEY_RANGE_ID
))
.map(|s| s.to_string()),
Some("0".to_string())
);
assert!(
headers
.get_optional_str(&HeaderName::from_static(
request_header_names::READ_FEED_KEY_TYPE
))
.is_none(),
"whole-PK-range targets must not emit x-ms-read-key-type"
);
assert!(headers
.get_optional_str(&HeaderName::from_static(request_header_names::START_EPK))
.is_none());
assert!(headers
.get_optional_str(&HeaderName::from_static(request_header_names::END_EPK))
.is_none());
}
#[test]
fn apply_headers_feed_range_emits_read_key_type_and_epk_bounds() {
let feed_range = FeedRange::new(
EffectivePartitionKey::from("10"),
EffectivePartitionKey::from("20"),
)
.unwrap();
let overrides = OperationOverrides {
partition_key_range_id: Some("pkrange".to_string()),
feed_range: Some(feed_range),
..Default::default()
};
let mut headers = azure_core::http::headers::Headers::new();
overrides
.apply_headers(&mut headers)
.expect("apply_headers should succeed");
assert_eq!(
headers
.get_optional_str(&HeaderName::from_static(
request_header_names::READ_FEED_KEY_TYPE
))
.map(|s| s.to_string()),
Some("EffectivePartitionKey".to_string())
);
assert_eq!(
headers
.get_optional_str(&HeaderName::from_static(request_header_names::START_EPK))
.map(|s| s.to_string()),
Some("10".to_string())
);
assert_eq!(
headers
.get_optional_str(&HeaderName::from_static(request_header_names::END_EPK))
.map(|s| s.to_string()),
Some("20".to_string())
);
}
#[test]
fn build_transport_request_feed_path_is_resolved() {
let operation = CosmosOperation::read_all_databases(test_account());
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert_eq!(request.url.path(), "/dbs");
}
#[test]
fn build_transport_request_single_resource_path_is_resolved() {
let db = DatabaseReference::from_name(test_account(), "mydb");
let operation = CosmosOperation::read_database(db);
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert_eq!(request.url.path(), "/dbs/mydb");
}
#[test]
fn build_transport_request_uses_operation_activity_id_when_present() {
let operation = CosmosOperation::read_all_databases(test_account())
.with_activity_id(ActivityId::from_string("operation-activity".to_string()));
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
let activity_header = request
.headers
.get_optional_str(&HeaderName::from_static("x-ms-activity-id"))
.expect("activity id should be set");
assert_eq!(activity_header, "operation-activity");
}
#[test]
fn build_transport_request_adds_partition_key_header_for_item_operation() {
let item_ref =
ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::read_item(item_ref);
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Retry,
deadline: Some(std::time::Instant::now() + Duration::from_secs(5)),
resolved_session_token: None,
throughput_control: None,
};
let overrides = OperationOverrides {
partition_key: Some(PartitionKey::from("pk1")),
..Default::default()
};
let request = build_transport_request(&operation, &overrides, None, &ctx)
.expect("request should build");
let partition_key_header = request
.headers
.get_optional_str(&HeaderName::from_static("x-ms-documentdb-partitionkey"))
.expect("partition key header should be set");
assert_eq!(partition_key_header, "[\"pk1\"]");
}
#[test]
fn build_transport_request_uses_routed_endpoint_url_directly() {
let operation =
CosmosOperation::read_database(DatabaseReference::from_name(test_account(), "mydb"));
let routing = RoutingDecision {
endpoint: CosmosEndpoint::regional_with_gateway20(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
Url::parse("https://test-westus2-thin.documents.azure.com:444/").unwrap(),
),
selected_url: Url::parse("https://test-westus2-thin.documents.azure.com:444/").unwrap(),
transport_mode: TransportMode::Gateway20,
};
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert_eq!(
request.url.as_str(),
"https://test-westus2-thin.documents.azure.com:444/dbs/mydb"
);
}
#[test]
fn build_transport_request_uses_default_url_for_global_endpoint() {
let operation =
CosmosOperation::read_database(DatabaseReference::from_name(test_account(), "mydb"));
let routing = RoutingDecision {
endpoint: CosmosEndpoint::global(
Url::parse("https://test.documents.azure.com:443/").unwrap(),
),
selected_url: Url::parse("https://test.documents.azure.com:443/").unwrap(),
transport_mode: TransportMode::Gateway,
};
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert_eq!(
request.url.as_str(),
"https://test.documents.azure.com/dbs/mydb"
);
}
#[test]
fn resolve_endpoint_uses_write_region_for_single_write_session_retry() {
let operation = CosmosOperation::read_all_databases(test_account());
let write_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let read_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![read_endpoint].into(),
preferred_write_endpoints: vec![write_endpoint.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: write_endpoint.clone(),
}));
let retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 1,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: Vec::new(),
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredWriteEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, write_endpoint);
}
#[test]
fn resolve_endpoint_deprioritizes_unavailable_over_global_fallback() {
let operation = CosmosOperation::read_all_databases(test_account());
let default_endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
let read_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
read_endpoint.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::TransportError,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![read_endpoint.clone()].into(),
preferred_write_endpoints: vec![default_endpoint.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: default_endpoint.clone(),
}));
let retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: Vec::new(),
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, read_endpoint);
}
#[test]
fn resolve_endpoint_ignores_write_forbidden_for_reads() {
let operation = CosmosOperation::read_all_databases(test_account());
let read_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
read_endpoint.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::WriteForbidden,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![read_endpoint.clone()].into(),
preferred_write_endpoints: vec![read_endpoint.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: read_endpoint.clone(),
}));
let retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: Vec::new(),
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, read_endpoint);
}
#[test]
fn stale_generation_advances_across_refreshed_endpoint_list() {
let operation = CosmosOperation::read_all_databases(test_account());
let endpoint_a = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let endpoint_b = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let endpoint_c = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 1,
preferred_read_endpoints: vec![
endpoint_a.clone(),
endpoint_b.clone(),
endpoint_c.clone(),
]
.into(),
preferred_write_endpoints: vec![
endpoint_a.clone(),
endpoint_b.clone(),
endpoint_c.clone(),
]
.into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: endpoint_a.clone(),
}));
let stale_retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0).next(3),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 3,
can_use_multiple_write_locations: true,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: Vec::new(),
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
let first_routing = super::resolve_endpoint(
&operation,
&stale_retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(first_routing.endpoint, endpoint_a);
let advanced_state = stale_retry_state
.advance_failover()
.advance_location(3, location.account.generation);
let second_routing = super::resolve_endpoint(
&operation,
&advanced_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(second_routing.endpoint, endpoint_b);
}
mod should_capture_session_token_from_status_tests {
use azure_core::http::StatusCode;
use crate::{
driver::pipeline::components::TransportOutcome,
models::{CosmosResponseHeaders, CosmosStatus, SubStatusCode},
};
use super::super::should_capture_session_token_from_status;
fn success_outcome() -> TransportOutcome {
TransportOutcome::Success {
status: CosmosStatus::new(StatusCode::Ok),
cosmos_headers: CosmosResponseHeaders::default(),
body: Vec::new(),
}
}
fn http_error_outcome(status: StatusCode) -> TransportOutcome {
TransportOutcome::HttpError {
status: CosmosStatus::new(status),
cosmos_headers: CosmosResponseHeaders::default(),
body: Vec::new(),
request_sent: crate::diagnostics::RequestSentStatus::Sent,
}
}
#[test]
fn captures_on_success() {
let outcome = success_outcome();
assert!(should_capture_session_token_from_status(None, &outcome));
}
#[test]
fn captures_on_409_conflict() {
let outcome = http_error_outcome(StatusCode::Conflict);
assert!(should_capture_session_token_from_status(None, &outcome));
}
#[test]
fn captures_on_412_precondition_failed() {
let outcome = http_error_outcome(StatusCode::PreconditionFailed);
assert!(should_capture_session_token_from_status(None, &outcome));
}
#[test]
fn skips_on_404_with_substatus_1002() {
let outcome = http_error_outcome(StatusCode::NotFound);
let substatus = SubStatusCode::READ_SESSION_NOT_AVAILABLE;
assert!(!should_capture_session_token_from_status(
Some(&substatus),
&outcome
));
}
#[test]
fn captures_on_404_without_substatus_1002() {
let outcome = http_error_outcome(StatusCode::NotFound);
assert!(should_capture_session_token_from_status(None, &outcome));
}
#[test]
fn skips_on_500_internal_server_error() {
let outcome = http_error_outcome(StatusCode::InternalServerError);
assert!(!should_capture_session_token_from_status(None, &outcome));
}
}
mod effective_consistency_tests {
use crate::{models::DefaultConsistencyLevel, options::ReadConsistencyStrategy};
#[test]
fn default_strategy_with_session_account() {
assert!(ReadConsistencyStrategy::Default
.is_session_effective(DefaultConsistencyLevel::Session));
}
#[test]
fn default_strategy_with_strong_account() {
assert!(!ReadConsistencyStrategy::Default
.is_session_effective(DefaultConsistencyLevel::Strong));
}
#[test]
fn session_strategy_overrides_account() {
assert!(ReadConsistencyStrategy::Session
.is_session_effective(DefaultConsistencyLevel::Strong));
}
#[test]
fn eventual_strategy_never_session() {
assert!(!ReadConsistencyStrategy::Eventual
.is_session_effective(DefaultConsistencyLevel::Session));
}
#[test]
fn consistent_prefix_not_session() {
assert!(!ReadConsistencyStrategy::Default
.is_session_effective(DefaultConsistencyLevel::ConsistentPrefix));
}
}
#[test]
fn resolve_endpoint_prefers_gateway20_for_dataplane_reads() {
let operation = CosmosOperation::read_item(ItemReference::from_name(
&test_container(),
PartitionKey::from("pk1"),
"doc1",
));
let endpoint = CosmosEndpoint::regional_with_gateway20(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
Url::parse("https://test-westus2-thin.documents.azure.com:444/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![endpoint.clone()].into(),
preferred_write_endpoints: vec![endpoint.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: endpoint.clone(),
}));
let retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
true,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, endpoint);
assert_eq!(routing.transport_mode, TransportMode::Gateway20);
assert_eq!(
routing.selected_url.as_str(),
"https://test-westus2-thin.documents.azure.com:444/"
);
}
#[test]
fn resolve_endpoint_skips_unavailable_region_when_gateway20_is_present() {
let operation = CosmosOperation::read_item(ItemReference::from_name(
&test_container(),
PartitionKey::from("pk1"),
"doc1",
));
let endpoint = CosmosEndpoint::regional_with_gateway20(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
Url::parse("https://test-westus2-thin.documents.azure.com:444/").unwrap(),
);
let fallback_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
endpoint.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::TransportError,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![endpoint.clone(), fallback_endpoint.clone()].into(),
preferred_write_endpoints: vec![endpoint].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: true,
default_endpoint: fallback_endpoint.clone(),
}));
let retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
3,
);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
true,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, fallback_endpoint);
}
#[test]
fn resolve_endpoint_metadata_falls_back_to_hub_when_all_excluded() {
let operation = CosmosOperation::read_all_databases(test_account());
let default_endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
let hub_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let read_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![read_endpoint.clone()].into(),
preferred_write_endpoints: vec![hub_endpoint.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: default_endpoint.clone(),
}));
let retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: vec!["westus2".into(), "eastus".into()],
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, hub_endpoint);
assert!(!routing.endpoint.is_global());
}
#[test]
fn resolve_endpoint_dataplane_excluded_read_region_routes_to_hub() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::read_item(item);
let default_endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
let hub_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let read_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![read_endpoint.clone()].into(),
preferred_write_endpoints: vec![hub_endpoint.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: default_endpoint.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: vec!["westus2".into()],
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
retry_state.is_dataplane = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, hub_endpoint);
assert!(!routing.endpoint.is_global());
}
#[test]
fn resolve_endpoint_dataplane_all_regions_excluded_falls_back_to_hub() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::read_item(item);
let default_endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
let hub_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let read_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![read_endpoint.clone()].into(),
preferred_write_endpoints: vec![hub_endpoint.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: default_endpoint.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: vec!["westus2".into(), "eastus".into()],
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
retry_state.is_dataplane = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, hub_endpoint);
assert!(!routing.endpoint.is_global());
}
#[test]
fn resolve_endpoint_dataplane_excluded_and_unavailable_still_uses_hub() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::read_item(item);
let default_endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
let hub_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let read_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
read_endpoint.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::ServiceUnavailable,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![read_endpoint.clone()].into(),
preferred_write_endpoints: vec![hub_endpoint.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: default_endpoint.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: vec!["westus2".into()],
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
retry_state.is_dataplane = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, hub_endpoint);
assert!(!routing.endpoint.is_global());
}
#[test]
fn resolve_endpoint_dataplane_isolated_region_unavailable_retries_same_region() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::read_item(item);
let default_endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
let hub_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let isolated_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
isolated_endpoint.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::ServiceUnavailable,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![isolated_endpoint.clone(), hub_endpoint.clone()].into(),
preferred_write_endpoints: vec![hub_endpoint.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: default_endpoint.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState {
location: LocationIndex::initial(0),
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 3,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: false,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: vec!["eastus".into()],
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
};
retry_state.is_dataplane = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, isolated_endpoint,
"region-isolation broken: resolve_endpoint silently fell back \
to a region the caller excluded"
);
assert_ne!(
routing.endpoint, hub_endpoint,
"region-isolation broken: resolve_endpoint chose the excluded hub region"
);
assert!(!routing.endpoint.is_global());
}
#[test]
fn resolve_endpoint_multi_write_isolated_region_unavailable_read_returns_isolated_not_excluded()
{
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::read_item(item);
let r1 = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let r2 = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
r2.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::ServiceUnavailable,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![r1.clone(), r2.clone()].into(),
preferred_write_endpoints: vec![r1.clone(), r2.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: true,
default_endpoint: r1.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
vec!["eastus".into()],
3,
2,
);
retry_state.is_dataplane = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, r2,
"isolation broken on read: resolve_endpoint did not return the \
caller's isolated region (r2)"
);
assert_ne!(
routing.endpoint, r1,
"isolation broken on read: resolve_endpoint silently routed to \
the excluded region (r1)"
);
}
#[test]
fn resolve_endpoint_multi_write_isolated_region_unavailable_write_retry_returns_isolated_not_excluded(
) {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let r1 = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let r2 = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
r2.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::ServiceUnavailable,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![r1.clone(), r2.clone()].into(),
preferred_write_endpoints: vec![r1.clone(), r2.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: true,
default_endpoint: r1.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
vec!["eastus".into()],
3,
2,
);
retry_state.is_dataplane = true;
retry_state
.pending_write_effects
.push(make_pending_partition_mark_for_region("westus2"));
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, r2,
"isolation broken on write retry: resolve_endpoint did not \
return the caller's isolated region (r2)"
);
assert_ne!(
routing.endpoint, r1,
"isolation broken on write retry: resolve_endpoint silently \
routed to the excluded region (r1) via the last-resort fallback"
);
}
#[test]
fn resolve_endpoint_single_write_excluded_hub_satellite_unavailable_read_returns_satellite() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::read_item(item);
let hub = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let satellite = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
satellite.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::ServiceUnavailable,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![hub.clone(), satellite.clone()].into(),
preferred_write_endpoints: vec![hub.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: hub.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
vec!["eastus".into()],
3,
2,
);
retry_state.is_dataplane = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, satellite,
"isolation broken on single-write read: resolve_endpoint did \
not return the caller's isolated region (satellite)"
);
assert_ne!(
routing.endpoint, hub,
"isolation broken on single-write read: resolve_endpoint \
silently routed to the excluded hub region"
);
}
#[test]
fn resolve_endpoint_picks_first_available_over_unavailable() {
let operation = CosmosOperation::read_all_databases(test_account());
let default_endpoint =
CosmosEndpoint::global(Url::parse("https://test.documents.azure.com:443/").unwrap());
let unavailable_endpoint = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let available_endpoint = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
unavailable_endpoint.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::ServiceUnavailable,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![
unavailable_endpoint.clone(),
available_endpoint.clone(),
]
.into(),
preferred_write_endpoints: vec![default_endpoint.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: default_endpoint.clone(),
}));
let retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, available_endpoint);
}
fn make_pending_partition_mark_for_region(
region: &'static str,
) -> crate::driver::routing::LocationEffect {
crate::driver::routing::LocationEffect::MarkPartitionUnavailable(
crate::driver::routing::UnavailablePartition {
partition_key_range_id: None,
region: Some(region.into()),
is_read: false,
is_partitioned_resource: true,
},
)
}
#[test]
fn resolve_endpoint_skips_in_flight_failed_write_region() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let west = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone(), west.clone()].into(),
preferred_write_endpoints: vec![east.clone(), west.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: east.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state
.pending_write_effects
.push(make_pending_partition_mark_for_region("eastus"));
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, west,
"in-flight skip set must route the retry to westus2"
);
}
#[test]
fn resolve_endpoint_honors_excluded_regions() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let west = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone(), west.clone()].into(),
preferred_write_endpoints: vec![east.clone(), west.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: east.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
vec!["westus2".into()],
3,
2,
);
retry_state.is_dataplane = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, east,
"excluded_regions must gate selection — got {:?}",
routing.endpoint
);
}
#[test]
fn resolve_endpoint_ppaf_write_uses_read_endpoints_as_primary_list() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let west = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![west.clone(), east.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
retry_state.ppaf_write_retry_allowed = true;
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, west,
"PPAF write must use the read endpoint list (preferred order) as the primary candidate set"
);
}
#[test]
fn resolve_endpoint_ppaf_falls_back_to_read_region_when_write_exhausted() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let west = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone(), west.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
retry_state.ppaf_write_retry_allowed = true;
retry_state
.pending_write_effects
.push(make_pending_partition_mark_for_region("eastus"));
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, west,
"PPAF write retry must fall back to a read region when all write regions are in the in-flight skip set"
);
}
#[test]
fn resolve_endpoint_ppaf_override_skipped_when_current_endpoint_failed_in_flight() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![north.clone(), central.clone()].into(),
preferred_write_endpoints: vec![north.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: north.clone(),
});
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let mut partitions = PartitionEndpointState::new(PartitionFailoverOptions::default());
partitions.per_partition_automatic_failover_enabled = true;
partitions.failover_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: central.clone(),
first_failed_endpoint: central.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 0,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
retry_state.ppaf_write_retry_allowed = true;
retry_state.partition_key_range_id = Some(pk_range_id);
retry_state
.pending_write_effects
.push(make_pending_partition_mark_for_region("centralus"));
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, north,
"PPAF override pointing at a region already in the in-flight skip set must be skipped, \
so cross-region retry can rotate to a different region instead of looping on the failed override"
);
}
#[test]
fn resolve_endpoint_ppaf_override_honored_when_current_endpoint_not_failed() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![north.clone(), central.clone()].into(),
preferred_write_endpoints: vec![north.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: north.clone(),
});
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let mut partitions = PartitionEndpointState::new(PartitionFailoverOptions::default());
partitions.per_partition_automatic_failover_enabled = true;
partitions.failover_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: central.clone(),
first_failed_endpoint: north.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 0,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
retry_state.ppaf_write_retry_allowed = true;
retry_state.partition_key_range_id = Some(pk_range_id);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, central,
"PPAF override with a healthy current_endpoint must be honored"
);
}
fn make_partition_state_with_ppcb_override(
pk_range_id: &super::PartitionKeyRangeId,
override_target: CosmosEndpoint,
) -> crate::driver::routing::partition_endpoint_state::PartitionEndpointState {
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let config = PartitionFailoverOptions::default();
let mut partitions = PartitionEndpointState::new(config);
partitions.per_partition_circuit_breaker_enabled = true;
partitions.circuit_breaker_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: override_target.clone(),
first_failed_endpoint: override_target,
failed_endpoints: Default::default(),
read_failure_count: 100,
write_failure_count: 100,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
partitions
}
#[test]
fn resolve_endpoint_ppcb_override_skipped_when_in_excluded_regions() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![north.clone(), central.clone()].into(),
preferred_write_endpoints: vec![north.clone(), central.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: north.clone(),
});
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let partitions = make_partition_state_with_ppcb_override(&pk_range_id, central.clone());
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state.partition_key_range_id = Some(pk_range_id);
retry_state.excluded_regions = vec![crate::options::Region::from("centralus")];
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, north,
"PPCB override pointing at a caller-excluded region must be skipped; \
fall-through to primary selection should pick the non-excluded region"
);
}
#[test]
fn resolve_endpoint_ppcb_override_skipped_when_current_endpoint_not_in_topology() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 1,
preferred_read_endpoints: vec![north.clone()].into(),
preferred_write_endpoints: vec![north.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: north.clone(),
});
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let partitions = make_partition_state_with_ppcb_override(&pk_range_id, central.clone());
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state.partition_key_range_id = Some(pk_range_id);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, north,
"PPCB override pointing at a region dropped from the post-refresh \
topology must be skipped; fall-through to primary selection should \
pick the only surviving region (`north`). The override map is not \
pruned by topology refresh, so `region_not_in_topology` in \
`ppcb_should_skip` is the only thing preventing a 1008 loop."
);
}
#[test]
fn resolve_endpoint_ppcb_override_skipped_when_endpoint_marked_unavailable() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
central.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::TransportError,
),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![north.clone(), central.clone()].into(),
preferred_write_endpoints: vec![north.clone(), central.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: true,
default_endpoint: north.clone(),
});
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let partitions = make_partition_state_with_ppcb_override(&pk_range_id, central.clone());
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state.partition_key_range_id = Some(pk_range_id);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, north,
"PPCB override pointing at an endpoint marked unavailable (e.g. transport-dead) \
must be skipped so the next attempt does not repeat the same connect failure"
);
}
#[test]
fn resolve_endpoint_ppcb_override_skipped_when_current_endpoint_failed_in_flight() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![north.clone(), central.clone()].into(),
preferred_write_endpoints: vec![north.clone(), central.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: north.clone(),
});
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let partitions = make_partition_state_with_ppcb_override(&pk_range_id, central.clone());
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state.partition_key_range_id = Some(pk_range_id);
retry_state
.pending_write_effects
.push(make_pending_partition_mark_for_region("centralus"));
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, north,
"PPCB override pointing at a region in the in-flight skip set must be skipped"
);
}
#[test]
fn resolve_endpoint_ppaf_override_honored_even_when_in_excluded_regions() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![north.clone(), central.clone()].into(),
preferred_write_endpoints: vec![north.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: north.clone(),
});
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let mut partitions = PartitionEndpointState::new(PartitionFailoverOptions::default());
partitions.per_partition_automatic_failover_enabled = true;
partitions.failover_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: central.clone(),
first_failed_endpoint: north.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 0,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
retry_state.ppaf_write_retry_allowed = true;
retry_state.partition_key_range_id = Some(pk_range_id);
retry_state.excluded_regions = vec![crate::options::Region::from("centralus")];
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, central,
"PPAF override target must be honored even when in excluded_regions: \
for a single-master account there is no other write region, so the SDK \
routes to the PPAF target and surfaces the failure to the caller"
);
}
#[test]
fn resolve_endpoint_ppaf_override_honored_even_when_endpoint_marked_unavailable() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let north = CosmosEndpoint::regional(
"northcentralus".into(),
Url::parse("https://test-northcentralus.documents.azure.com:443/").unwrap(),
);
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
central.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::TransportError,
),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![north.clone(), central.clone()].into(),
preferred_write_endpoints: vec![north.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: north.clone(),
});
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let mut partitions = PartitionEndpointState::new(PartitionFailoverOptions::default());
partitions.per_partition_automatic_failover_enabled = true;
partitions.failover_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: central.clone(),
first_failed_endpoint: north.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 0,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
retry_state.ppaf_write_retry_allowed = true;
retry_state.partition_key_range_id = Some(pk_range_id);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, central,
"PPAF override target must be honored even when marked unavailable: \
single-master account has no other write region to fail over to"
);
}
#[test]
fn resolve_endpoint_ppcb_override_probe_skipped_when_first_failed_endpoint_failed_in_flight() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let east = CosmosEndpoint::regional(
"eastus2".into(),
Url::parse("https://test-eastus2.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![central.clone(), east.clone()].into(),
preferred_write_endpoints: vec![central.clone(), east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: central.clone(),
});
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let mut partitions = PartitionEndpointState::new(PartitionFailoverOptions::default());
partitions.per_partition_circuit_breaker_enabled = true;
partitions.circuit_breaker_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: east.clone(),
first_failed_endpoint: central.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 0,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::ProbeCandidate,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true, Vec::new(),
3,
2,
);
retry_state.partition_key_range_id = Some(pk_range_id);
retry_state
.pending_write_effects
.push(make_pending_partition_mark_for_region("centralus"));
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_ne!(
routing.endpoint, central,
"PPCB ProbeCandidate must skip the probe target when its region is already in \
the in-flight skip set; otherwise retry pins to the failing probe region"
);
}
#[test]
fn resolve_endpoint_ppcb_override_honored_when_current_endpoint_not_failed() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let east = CosmosEndpoint::regional(
"eastus2".into(),
Url::parse("https://test-eastus2.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![central.clone(), east.clone()].into(),
preferred_write_endpoints: vec![central.clone(), east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: central.clone(),
});
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let mut partitions = PartitionEndpointState::new(PartitionFailoverOptions::default());
partitions.per_partition_circuit_breaker_enabled = true;
let write_threshold = partitions.config.write_failure_threshold();
partitions.circuit_breaker_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: east.clone(),
first_failed_endpoint: central.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: write_threshold as i32 + 10,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state.partition_key_range_id = Some(pk_range_id);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, east,
"PPCB override with a healthy current_endpoint must be honored on a fresh attempt"
);
}
#[test]
fn resolve_endpoint_ppcb_override_pinned_to_unavailable_endpoint_reproduces_prod_403_3_loop() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let east = CosmosEndpoint::regional(
"eastus2".into(),
Url::parse("https://test-eastus2.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
central.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::WriteForbidden,
),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![central.clone(), east.clone()].into(),
preferred_write_endpoints: vec![central.clone(), east.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: true,
default_endpoint: central.clone(),
});
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::options::PartitionFailoverOptions;
let pk_range_id: super::PartitionKeyRangeId = "0".parse().unwrap();
let mut partitions = PartitionEndpointState::new(PartitionFailoverOptions::default());
partitions.per_partition_circuit_breaker_enabled = true;
partitions.circuit_breaker_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: central.clone(),
first_failed_endpoint: central.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 10,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(account, Arc::new(partitions));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state.partition_key_range_id = Some(pk_range_id);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_ne!(
routing.endpoint, central,
"BUG: PPCB override pinned the partition to centralus even though \
centralus is in account.unavailable_endpoints. This is the \
production 4-attempt-all-to-central failure mode. The PPCB \
override-skip path in resolve_endpoint must honor \
account.unavailable_endpoints (and excluded_regions) the same \
way try_select_endpoint does."
);
assert_eq!(
routing.endpoint, east,
"with central in unavailable_endpoints and east the only other write \
region, resolve_endpoint must route the next attempt to east"
);
}
#[test]
fn resolve_endpoint_rotates_to_next_region_via_location_index_on_multi_write() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let central = CosmosEndpoint::regional(
"centralus".into(),
Url::parse("https://test-centralus.documents.azure.com:443/").unwrap(),
);
let east = CosmosEndpoint::regional(
"eastus2".into(),
Url::parse("https://test-eastus2.documents.azure.com:443/").unwrap(),
);
let account = Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![central.clone(), east.clone()].into(),
preferred_write_endpoints: vec![central.clone(), east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: true,
default_endpoint: central.clone(),
});
let location = LocationSnapshot::for_tests(account);
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
true,
Vec::new(),
3,
2,
);
retry_state.location = retry_state.location.next_for_generation(2, 0);
assert_eq!(
retry_state.location.index(),
1,
"precondition: LocationIndex must have advanced from 0 to 1"
);
assert!(
retry_state.pending_write_effects.is_empty(),
"precondition: multi-write deferred bucket must be empty"
);
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(
routing.endpoint, east,
"BUG: with LocationIndex=1 and empty in_flight_failed, \
resolve_endpoint must route to east. The multi-write \
PPCB-managed path has no other skip signal — if LocationIndex \
rotation regresses, the next attempt will pin to the just-failed \
centralus and reproduce the same-region-retry loop."
);
}
#[test]
fn resolve_endpoint_does_not_skip_for_reads() {
let operation = CosmosOperation::read_all_databases(test_account());
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}));
let mut retry_state = crate::driver::pipeline::components::OperationRetryState::initial(
0,
false,
Vec::new(),
3,
2,
);
retry_state
.pending_write_effects
.push(make_pending_partition_mark_for_region("eastus"));
let routing = super::resolve_endpoint(
&operation,
&retry_state,
&location,
false,
Duration::from_secs(60),
);
assert_eq!(routing.endpoint, east);
}
#[test]
fn already_applied_skips_endpoint_mark_when_endpoint_in_unavailable_set() {
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let mut unavailable = std::collections::HashMap::new();
unavailable.insert(
east.url().clone(),
(
std::time::Instant::now(),
crate::driver::routing::UnavailableReason::TransportError,
),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: unavailable,
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}));
let effect = LocationEffect::MarkEndpointUnavailable {
endpoint: east,
reason: crate::driver::routing::UnavailableReason::TransportError,
};
assert!(
super::is_effect_already_applied(&effect, &location),
"endpoint already in unavailable set must be considered already applied"
);
}
#[test]
fn already_applied_returns_false_when_endpoint_not_in_unavailable_set() {
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}));
let effect = LocationEffect::MarkEndpointUnavailable {
endpoint: east,
reason: crate::driver::routing::UnavailableReason::TransportError,
};
assert!(!super::is_effect_already_applied(&effect, &location));
}
#[test]
fn already_applied_skips_partition_mark_when_override_already_moved_past_failed_region() {
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::driver::routing::partition_key_range_id::PartitionKeyRangeId;
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let west = CosmosEndpoint::regional(
"westus2".into(),
Url::parse("https://test-westus2.documents.azure.com:443/").unwrap(),
);
let pk_range_id = PartitionKeyRangeId::from(String::from("0"));
let mut partitions = PartitionEndpointState {
per_partition_automatic_failover_enabled: true,
..Default::default()
};
partitions.failover_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: west.clone(), first_failed_endpoint: east.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 0,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(
Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone(), west.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}),
Arc::new(partitions),
);
let effect = LocationEffect::MarkPartitionUnavailable(
crate::driver::routing::UnavailablePartition {
partition_key_range_id: Some(pk_range_id),
region: Some("eastus".into()),
is_read: false,
is_partitioned_resource: true,
},
);
assert!(
super::is_effect_already_applied(&effect, &location),
"PPAF override that has already moved past the failed region must be considered already applied"
);
}
#[test]
fn already_applied_returns_false_when_partition_override_still_on_failed_region() {
use crate::driver::routing::partition_endpoint_state::{
HealthStatus, PartitionEndpointState, PartitionFailoverEntry,
};
use crate::driver::routing::partition_key_range_id::PartitionKeyRangeId;
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let pk_range_id = PartitionKeyRangeId::from(String::from("0"));
let mut partitions = PartitionEndpointState {
per_partition_automatic_failover_enabled: true,
..Default::default()
};
partitions.failover_overrides.insert(
pk_range_id.clone(),
PartitionFailoverEntry {
current_endpoint: east.clone(), first_failed_endpoint: east.clone(),
failed_endpoints: Default::default(),
read_failure_count: 0,
write_failure_count: 0,
first_failure_time: std::time::Instant::now(),
last_failure_time: std::time::Instant::now(),
health_status: HealthStatus::Unhealthy,
failback_jitter: Duration::ZERO,
},
);
let location = LocationSnapshot::for_tests_with_partitions(
Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}),
Arc::new(partitions),
);
let effect = LocationEffect::MarkPartitionUnavailable(
crate::driver::routing::UnavailablePartition {
partition_key_range_id: Some(pk_range_id),
region: Some("eastus".into()),
is_read: false,
is_partitioned_resource: true,
},
);
assert!(
!super::is_effect_already_applied(&effect, &location),
"override still pointing at the failed region must NOT be skipped"
);
}
#[test]
fn already_applied_returns_false_when_no_partition_override_exists() {
let east = CosmosEndpoint::regional(
"eastus".into(),
Url::parse("https://test-eastus.documents.azure.com:443/").unwrap(),
);
let location = LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: vec![east.clone()].into(),
preferred_write_endpoints: vec![east.clone()].into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: east.clone(),
}));
let effect = LocationEffect::MarkPartitionUnavailable(
crate::driver::routing::UnavailablePartition {
partition_key_range_id: Some(
crate::driver::routing::partition_key_range_id::PartitionKeyRangeId::from(
String::from("0"),
),
),
region: Some("eastus".into()),
is_read: false,
is_partitioned_resource: true,
},
);
assert!(!super::is_effect_already_applied(&effect, &location));
}
#[test]
fn build_transport_request_sets_is_upsert_header() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::upsert_item(item).with_body(b"{}".to_vec());
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
let is_upsert = request
.headers
.get_optional_str(&HeaderName::from_static("x-ms-documentdb-is-upsert"))
.expect("is-upsert header should be set");
assert_eq!(is_upsert, "true");
assert_eq!(
request.url.path(),
"/dbs/testdb/colls/testcontainer/docs",
"upsert should POST to the collection feed, not /docs/doc1"
);
}
#[test]
fn build_transport_request_omits_is_upsert_header_for_create() {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert!(
request
.headers
.get_optional_str(&HeaderName::from_static("x-ms-documentdb-is-upsert"))
.is_none(),
"is-upsert header should not be set for create"
);
assert_eq!(
request.url.path(),
"/dbs/testdb/colls/testcontainer/docs",
"create should POST to the collection feed, not /docs/doc1"
);
}
#[test]
fn build_transport_request_sets_batch_headers() {
let operation = CosmosOperation::batch(test_container(), PartitionKey::from("pk1"))
.with_body(b"[]".to_vec());
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert_eq!(
request
.headers
.get_optional_str(&HeaderName::from_static("x-ms-cosmos-is-batch-request")),
Some("True"),
"is-batch-request header should be set"
);
assert_eq!(
request
.headers
.get_optional_str(&HeaderName::from_static("x-ms-cosmos-batch-atomic")),
Some("True"),
"batch-atomic header should be set"
);
assert_eq!(
request.headers.get_optional_str(&HeaderName::from_static(
"x-ms-cosmos-batch-continue-on-error"
)),
Some("False"),
"batch-continue-on-error header should be set"
);
}
#[test]
fn build_transport_request_omits_batch_headers_for_create() {
let container = test_container();
let operation = CosmosOperation::create_item(ItemReference::from_name(
&container,
PartitionKey::from("pk1"),
"doc1",
))
.with_body(b"{}".to_vec());
let routing = test_routing();
let activity_id = ActivityId::from_string("default-activity".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert!(
request
.headers
.get_optional_str(&HeaderName::from_static("x-ms-cosmos-is-batch-request"))
.is_none(),
"batch headers should not be set for create"
);
}
#[test]
fn build_transport_request_sets_priority_level_header() {
let container = test_container();
let operation = CosmosOperation::read_item(ItemReference::from_name(
&container,
PartitionKey::from("pk1"),
"doc1",
));
let routing = test_routing();
let activity_id = ActivityId::new_uuid();
let throughput_control = Some(ResolvedThroughputControl {
throughput_bucket: None,
priority_level: Some(PriorityLevel::Low),
});
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.unwrap();
let priority = request
.headers
.get_optional_str(&HeaderName::from_static(
request_header_names::PRIORITY_LEVEL,
))
.expect("priority level header should be set");
assert_eq!(priority, "Low");
assert!(request
.headers
.get_optional_str(&HeaderName::from_static(
request_header_names::THROUGHPUT_BUCKET
))
.is_none());
}
#[test]
fn build_transport_request_sets_throughput_bucket_header() {
let container = test_container();
let operation = CosmosOperation::read_item(ItemReference::from_name(
&container,
PartitionKey::from("pk1"),
"doc1",
));
let routing = test_routing();
let activity_id = ActivityId::new_uuid();
let throughput_control = Some(ResolvedThroughputControl {
throughput_bucket: Some(42),
priority_level: None,
});
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.unwrap();
let bucket = request
.headers
.get_optional_str(&HeaderName::from_static(
request_header_names::THROUGHPUT_BUCKET,
))
.expect("throughput bucket header should be set");
assert_eq!(bucket, "42");
assert!(request
.headers
.get_optional_str(&HeaderName::from_static(
request_header_names::PRIORITY_LEVEL
))
.is_none());
}
#[test]
fn build_transport_request_sets_both_throughput_headers() {
let container = test_container();
let operation = CosmosOperation::read_item(ItemReference::from_name(
&container,
PartitionKey::from("pk1"),
"doc1",
));
let routing = test_routing();
let activity_id = ActivityId::new_uuid();
let throughput_control = Some(ResolvedThroughputControl {
throughput_bucket: Some(100),
priority_level: Some(PriorityLevel::High),
});
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control,
};
let request =
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.unwrap();
assert_eq!(
request.headers.get_optional_str(&HeaderName::from_static(
request_header_names::PRIORITY_LEVEL
)),
Some("High")
);
assert_eq!(
request.headers.get_optional_str(&HeaderName::from_static(
request_header_names::THROUGHPUT_BUCKET
)),
Some("100")
);
}
#[test]
fn build_transport_request_auto_emits_query_headers_for_query_operations() {
let container = test_container();
let pk_def = container.partition_key_definition().clone();
let feed_range = FeedRange::for_partition(PartitionKey::from("pk1"), &pk_def);
let op = CosmosOperation::query_items(container.clone(), Some(feed_range))
.with_body(br#"{"query":"SELECT * FROM c"}"#.to_vec());
assert_query_headers_present(&op, "query_items (single partition)");
let op = CosmosOperation::query_items(container, Some(FeedRange::full()))
.with_body(br#"{"query":"SELECT * FROM c"}"#.to_vec());
assert_query_headers_present(&op, "query_items (cross partition)");
let op = CosmosOperation::query_offers(test_account())
.with_body(br#"{"query":"SELECT * FROM root"}"#.to_vec());
assert_query_headers_present(&op, "query_offers");
}
fn assert_query_headers_present(op: &CosmosOperation, label: &str) {
let routing = test_routing();
let activity_id = ActivityId::new_uuid();
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
let req = build_transport_request(op, &OperationOverrides::default(), None, &ctx)
.expect("request should build");
assert_eq!(
req.headers
.get_optional_str(&HeaderName::from_static(request_header_names::IS_QUERY)),
Some("True"),
"{label}: x-ms-documentdb-isquery should be 'True'"
);
assert_eq!(
req.headers
.get_optional_str(&azure_core::http::headers::CONTENT_TYPE),
Some(crate::models::cosmos_headers::QUERY_CONTENT_TYPE),
"{label}: Content-Type should be application/query+json"
);
}
fn retry_state_with_counts(
failover_retry_count: u32,
session_token_retry_count: u32,
) -> super::OperationRetryState {
let mut state = super::OperationRetryState::initial(0, false, Vec::new(), 3, 1);
state.failover_retry_count = failover_retry_count;
state.session_token_retry_count = session_token_retry_count;
state
}
#[test]
fn execution_context_initial_when_no_retries() {
let state = retry_state_with_counts(0, 0);
assert!(matches!(
super::compute_execution_context(&state),
ExecutionContext::Initial
));
}
#[test]
fn execution_context_retry_when_session_retry_active() {
let state = retry_state_with_counts(1, 1);
assert!(matches!(
super::compute_execution_context(&state),
ExecutionContext::Retry
));
let state = retry_state_with_counts(0, 1);
assert!(matches!(
super::compute_execution_context(&state),
ExecutionContext::Retry
));
}
#[test]
fn execution_context_region_failover_when_only_failover_active() {
let state = retry_state_with_counts(1, 0);
assert!(matches!(
super::compute_execution_context(&state),
ExecutionContext::RegionFailover
));
}
fn build_minimal_transport_request() -> super::TransportRequest {
let operation = CosmosOperation::read_all_databases(test_account());
let routing = test_routing();
let activity_id = ActivityId::from_string("hub-region-test".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build")
}
#[test]
fn transport_request_emits_hub_region_header_when_latched() {
let mut request = build_minimal_transport_request();
let mut state = super::OperationRetryState::initial(0, false, Vec::new(), 3, 1);
state.is_dataplane = true;
state.hub_region_processing_only = true;
super::apply_hub_region_header(&mut request, &state);
let value = request.headers.get_optional_str(&HeaderName::from_static(
request_header_names::HUB_REGION_PROCESSING_ONLY,
));
assert_eq!(value, Some("True"));
}
#[test]
fn transport_request_omits_hub_region_header_when_not_latched() {
let mut request = build_minimal_transport_request();
let state = super::OperationRetryState::initial(0, false, Vec::new(), 3, 1);
assert!(!state.hub_region_processing_only);
super::apply_hub_region_header(&mut request, &state);
let value = request.headers.get_optional_str(&HeaderName::from_static(
request_header_names::HUB_REGION_PROCESSING_ONLY,
));
assert!(
value.is_none(),
"hub-region header must not be present when latch is unset, got {value:?}",
);
}
fn build_write_transport_request() -> super::TransportRequest {
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
let routing = test_routing();
let activity_id = ActivityId::from_string("tentative-writes-test".to_string());
let ctx = TransportRequestContext {
routing: &routing,
activity_id: &activity_id,
execution_context: ExecutionContext::Initial,
deadline: None,
resolved_session_token: None,
throughput_control: None,
};
build_transport_request(&operation, &OperationOverrides::default(), None, &ctx)
.expect("request should build")
}
fn tentative_writes_header(request: &super::TransportRequest) -> Option<&str> {
request.headers.get_optional_str(&HeaderName::from_static(
request_header_names::ALLOW_TENTATIVE_WRITES,
))
}
#[test]
fn tentative_writes_header_set_for_write_on_multi_write_account() {
let mut request = build_write_transport_request();
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
super::apply_tentative_writes_header(&mut request, &operation, true);
assert_eq!(tentative_writes_header(&request), Some("true"));
}
#[test]
fn tentative_writes_header_omitted_for_read_on_multi_write_account() {
let mut request = build_minimal_transport_request();
let operation = CosmosOperation::read_all_databases(test_account());
super::apply_tentative_writes_header(&mut request, &operation, true);
assert!(tentative_writes_header(&request).is_none());
}
#[test]
fn tentative_writes_header_omitted_for_write_on_single_master_account() {
let mut request = build_write_transport_request();
let item = ItemReference::from_name(&test_container(), PartitionKey::from("pk1"), "doc1");
let operation = CosmosOperation::create_item(item).with_body(b"{}".to_vec());
super::apply_tentative_writes_header(&mut request, &operation, false);
assert!(tentative_writes_header(&request).is_none());
}
#[tokio::test]
async fn failover_delay_none_returns_immediately() {
let start = std::time::Instant::now();
super::apply_failover_delay(None).await;
assert!(start.elapsed() < Duration::from_millis(50));
}
#[tokio::test]
async fn failover_delay_zero_returns_immediately() {
let start = std::time::Instant::now();
super::apply_failover_delay(Some(Duration::ZERO)).await;
assert!(start.elapsed() < Duration::from_millis(50));
}
#[tokio::test(start_paused = true)]
async fn failover_delay_real_value_actually_sleeps() {
let start = tokio::time::Instant::now();
super::apply_failover_delay(Some(Duration::from_secs(5))).await;
assert!(start.elapsed() >= Duration::from_secs(5));
}
fn empty_options_view() -> crate::options::OperationOptionsView<'static> {
crate::options::OperationOptionsView::new(None, None, None, None)
}
fn test_diagnostics() -> crate::diagnostics::DiagnosticsContextBuilder {
crate::diagnostics::DiagnosticsContextBuilder::new(
crate::models::ActivityId::from_string("test-deadline".to_owned()),
std::sync::Arc::new(crate::options::DiagnosticsOptions::default()),
)
}
#[test]
fn enforce_deadline_none_is_ok() {
let options = empty_options_view();
let diagnostics = test_diagnostics();
let result = super::enforce_deadline_or_timeout(None, &options, diagnostics);
assert!(result.is_ok());
}
#[test]
fn enforce_deadline_in_future_is_ok() {
let options = empty_options_view();
let diagnostics = test_diagnostics();
let deadline = std::time::Instant::now() + Duration::from_secs(60);
let result = super::enforce_deadline_or_timeout(Some(deadline), &options, diagnostics);
assert!(result.is_ok());
}
#[test]
fn enforce_deadline_in_past_returns_timeout_error_with_diagnostics() {
let options = empty_options_view();
let diagnostics = test_diagnostics();
let deadline = std::time::Instant::now() - Duration::from_millis(1);
let result = super::enforce_deadline_or_timeout(Some(deadline), &options, diagnostics);
let err = result.expect_err("past deadline should produce an error");
let msg = err.to_string();
assert!(
msg.contains("end-to-end operation timeout exceeded"),
"unexpected error message: {msg}"
);
assert!(
err.diagnostics().is_some(),
"timeout error must carry finalized diagnostics"
);
}
fn http_result(status_code: u16, sub_status: Option<u32>) -> super::TransportResult {
use azure_core::http::StatusCode;
let mut status = crate::models::CosmosStatus::new(StatusCode::from(status_code));
if let Some(v) = sub_status {
status = status.with_sub_status(v as u16);
}
super::TransportResult::from_http_response(
status,
crate::models::CosmosResponseHeaders::default(),
Vec::new(),
)
}
#[test]
fn classify_hedge_result_success_is_final() {
let tr = http_result(200, None);
assert!(matches!(
super::classify_hedge_result(Ok(tr)),
super::HedgeClass::Final(_)
));
}
#[test]
fn classify_hedge_result_409_conflict_is_final() {
let tr = http_result(409, None);
assert!(matches!(
super::classify_hedge_result(Ok(tr)),
super::HedgeClass::Final(_)
));
}
#[test]
fn classify_hedge_result_503_is_transient() {
let tr = http_result(503, None);
assert!(matches!(
super::classify_hedge_result(Ok(tr)),
super::HedgeClass::Transient
));
}
#[test]
fn classify_hedge_result_404_1002_is_transient() {
let tr = http_result(404, Some(1002));
assert!(matches!(
super::classify_hedge_result(Ok(tr)),
super::HedgeClass::Transient
));
}
#[test]
fn classify_hedge_result_deadline_exceeded_is_transient() {
let tr =
super::TransportResult::deadline_exceeded(crate::diagnostics::RequestSentStatus::Sent);
assert!(matches!(
super::classify_hedge_result(Ok(tr)),
super::HedgeClass::Transient
));
}
#[test]
fn classify_hedge_result_request_build_error_is_transient() {
let err = crate::error::CosmosError::builder()
.with_message("synthetic build error")
.build();
assert!(matches!(
super::classify_hedge_result(Err(err)),
super::HedgeClass::Transient
));
}
#[test]
fn result_is_final_success_is_true() {
let tr = http_result(200, None);
assert!(super::result_is_final(&tr));
}
#[test]
fn result_is_final_409_conflict_is_true() {
let tr = http_result(409, None);
assert!(super::result_is_final(&tr));
}
#[test]
fn result_is_final_503_is_false() {
let tr = http_result(503, None);
assert!(!super::result_is_final(&tr));
}
#[test]
fn result_is_final_404_1002_is_false() {
let tr = http_result(404, Some(1002));
assert!(!super::result_is_final(&tr));
}
#[test]
fn result_is_final_deadline_exceeded_is_false() {
let tr =
super::TransportResult::deadline_exceeded(crate::diagnostics::RequestSentStatus::Sent);
assert!(!super::result_is_final(&tr));
}
#[test]
fn result_is_final_429_ru_budget_and_hot_partition_are_final() {
assert!(super::result_is_final(&http_result(429, Some(3200)))); assert!(super::result_is_final(&http_result(429, Some(3210)))); assert!(super::result_is_final(&http_result(429, Some(3214)))); }
#[test]
fn result_is_final_429_generic_and_3092_are_transient() {
assert!(!super::result_is_final(&http_result(429, None)));
assert!(!super::result_is_final(&http_result(429, Some(3092))));
}
#[test]
fn result_is_final_agrees_with_classify_hedge_result() {
let cases = [
http_result(200, None),
http_result(201, None),
http_result(204, None),
http_result(400, None),
http_result(404, None),
http_result(404, Some(1002)),
http_result(409, None),
http_result(412, None),
http_result(429, None),
http_result(429, Some(3092)),
http_result(429, Some(3200)),
http_result(429, Some(3210)),
http_result(429, Some(3214)),
http_result(500, None),
http_result(503, None),
super::TransportResult::deadline_exceeded(crate::diagnostics::RequestSentStatus::Sent),
];
for tr in cases {
let by_peek = super::result_is_final(&tr);
let by_classify = matches!(
super::classify_hedge_result(Ok(tr)),
super::HedgeClass::Final(_)
);
assert_eq!(
by_peek, by_classify,
"result_is_final must agree with classify_hedge_result",
);
}
}
#[test]
fn finalize_hedge_attempt_http_error_returns_error_with_status() {
let tr = Box::new(http_result(409, None));
let diagnostics = test_diagnostics();
let err = super::finalize_hedge_attempt(tr, diagnostics)
.expect_err("409 should be surfaced as an error");
assert_eq!(u16::from(err.status().status_code()), 409);
}
#[test]
fn finalize_hedge_attempt_deadline_returns_other_error() {
let tr = Box::new(super::TransportResult::deadline_exceeded(
crate::diagnostics::RequestSentStatus::Sent,
));
let diagnostics = test_diagnostics();
let err = super::finalize_hedge_attempt(tr, diagnostics)
.expect_err("deadline should produce an error");
assert!(err.to_string().contains("deadline exceeded"));
}
#[test]
fn diagnostics_clone_for_hedge_attempt_starts_empty() {
let parent = test_diagnostics();
let child = parent.clone_for_hedge_attempt();
assert_eq!(child.request_count(), 0);
}
#[test]
fn diagnostics_merge_hedge_attempt_absorbs_requests() {
use crate::diagnostics::{TransportHttpVersion, TransportKind};
let mut parent = test_diagnostics();
let mut child = parent.clone_for_hedge_attempt();
let endpoint = crate::driver::routing::CosmosEndpoint::global(
url::Url::parse("https://acct.example/").unwrap(),
);
let _ = child.start_request(
super::ExecutionContext::Hedging,
super::PipelineType::DataPlane,
super::TransportSecurity::Secure,
TransportKind::Gateway,
TransportHttpVersion::Http11,
&endpoint,
);
assert_eq!(child.request_count(), 1);
assert_eq!(parent.request_count(), 0);
parent.merge_hedge_attempt(child);
assert_eq!(parent.request_count(), 1);
}
#[test]
fn deadline_elapsed_none_is_false() {
assert!(!super::deadline_elapsed(None));
}
#[test]
fn deadline_elapsed_future_is_false() {
let d = std::time::Instant::now() + Duration::from_secs(60);
assert!(!super::deadline_elapsed(Some(d)));
}
#[test]
fn deadline_elapsed_past_is_true() {
let d = std::time::Instant::now() - Duration::from_millis(1);
assert!(super::deadline_elapsed(Some(d)));
}
#[test]
fn application_cancelled_error_carries_app_cancel_message() {
let err = super::application_cancelled_error(test_diagnostics());
let msg = err.to_string();
assert!(
msg.contains("cancelled by application deadline"),
"unexpected error message: {msg}"
);
let status = err.status();
assert_eq!(
status.status_code(),
azure_core::http::StatusCode::RequestTimeout
);
assert_eq!(
status.sub_status(),
Some(crate::models::SubStatusCode::CLIENT_OPERATION_TIMEOUT)
);
assert!(
err.diagnostics().is_some(),
"application_cancelled_error must graft diagnostics"
);
}
#[test]
fn harvest_window_is_50ms() {
assert_eq!(super::HARVEST_WINDOW, Duration::from_millis(50));
}
#[tokio::test]
async fn deadline_signal_none_does_not_complete() {
let fut = super::deadline_signal(None);
let timer = Box::pin(azure_core::sleep(
azure_core::time::Duration::try_from(Duration::from_millis(20)).unwrap(),
));
match futures::future::select(fut, timer).await {
futures::future::Either::Right(((), _)) => { }
futures::future::Either::Left(((), _)) => {
panic!("deadline_signal(None) must never resolve");
}
}
}
#[tokio::test]
async fn deadline_signal_past_resolves_immediately() {
let past = std::time::Instant::now() - Duration::from_millis(10);
let fut = super::deadline_signal(Some(past));
let timer = Box::pin(azure_core::sleep(
azure_core::time::Duration::try_from(Duration::from_millis(50)).unwrap(),
));
match futures::future::select(fut, timer).await {
futures::future::Either::Left(((), _)) => { }
futures::future::Either::Right(((), _)) => {
panic!("deadline_signal(past) must resolve before a 50ms sleep");
}
}
}
#[tokio::test]
async fn harvest_remaining_attempt_merges_diagnostics_when_attempt_completes_in_window() {
use crate::diagnostics::{TransportHttpVersion, TransportKind};
let mut parent = test_diagnostics();
let mut child = parent.clone_for_hedge_attempt();
let endpoint = crate::driver::routing::CosmosEndpoint::global(
url::Url::parse("https://acct.example/").unwrap(),
);
let _ = child.start_request(
super::ExecutionContext::Hedging,
super::PipelineType::DataPlane,
super::TransportSecurity::Secure,
TransportKind::Gateway,
TransportHttpVersion::Http11,
&endpoint,
);
let attempt = Box::pin(async move {
(
Err::<super::TransportResult, _>(
crate::error::CosmosError::builder()
.with_message("synthetic transport error")
.build(),
),
child,
)
});
super::harvest_remaining_attempt(attempt, &mut parent, super::HARVEST_WINDOW).await;
assert_eq!(parent.request_count(), 1);
}
#[tokio::test(start_paused = true)]
async fn harvest_remaining_attempt_drops_attempt_when_window_exceeded() {
use crate::diagnostics::{TransportHttpVersion, TransportKind};
let mut parent = test_diagnostics();
let mut child = parent.clone_for_hedge_attempt();
let endpoint = crate::driver::routing::CosmosEndpoint::global(
url::Url::parse("https://acct.example/").unwrap(),
);
let _ = child.start_request(
super::ExecutionContext::Hedging,
super::PipelineType::DataPlane,
super::TransportSecurity::Secure,
TransportKind::Gateway,
TransportHttpVersion::Http11,
&endpoint,
);
let attempt = Box::pin(async move {
azure_core::sleep(
azure_core::time::Duration::try_from(Duration::from_secs(60)).unwrap(),
)
.await;
(
Err::<super::TransportResult, _>(
crate::error::CosmosError::builder()
.with_message("should not reach here")
.build(),
),
child,
)
});
let parent_before = parent.request_count();
super::harvest_remaining_attempt(attempt, &mut parent, super::HARVEST_WINDOW).await;
assert_eq!(parent.request_count(), parent_before);
}
#[test]
fn shared_hub_region_latch_eligibility_dataplane_single_master() {
assert!(super::should_build_shared_hub_region_latch(
super::PipelineType::DataPlane,
false, ));
}
#[test]
fn shared_hub_region_latch_eligibility_skip_multi_master() {
assert!(!super::should_build_shared_hub_region_latch(
super::PipelineType::DataPlane,
true, ));
}
#[test]
fn shared_hub_region_latch_eligibility_skip_metadata() {
assert!(!super::should_build_shared_hub_region_latch(
super::PipelineType::Metadata,
false,
));
assert!(!super::should_build_shared_hub_region_latch(
super::PipelineType::Metadata,
true,
));
}
#[test]
fn should_emit_hub_region_header_per_state_only() {
assert!(super::should_emit_hub_region_header(true, None));
}
#[test]
fn should_emit_hub_region_header_shared_only() {
use std::sync::{
atomic::{AtomicBool, Ordering},
Arc,
};
let shared = Arc::new(AtomicBool::new(true));
assert!(super::should_emit_hub_region_header(false, Some(&shared)));
assert!(shared.load(Ordering::Acquire));
}
#[test]
fn should_emit_hub_region_header_neither_latched() {
use std::sync::{atomic::AtomicBool, Arc};
let shared = Arc::new(AtomicBool::new(false));
assert!(!super::should_emit_hub_region_header(false, None));
assert!(!super::should_emit_hub_region_header(false, Some(&shared)));
}
#[test]
fn apply_hub_region_header_emits_when_only_shared_latch_set() {
use std::sync::{atomic::AtomicBool, Arc};
let mut request = build_minimal_transport_request();
let shared = Arc::new(AtomicBool::new(true));
let mut state = super::OperationRetryState::initial(0, false, Vec::new(), 3, 1);
state.is_dataplane = true;
state = state.with_shared_hub_region_latch(shared);
assert!(!state.hub_region_processing_only);
super::apply_hub_region_header(&mut request, &state);
let value = request.headers.get_optional_str(&HeaderName::from_static(
request_header_names::HUB_REGION_PROCESSING_ONLY,
));
assert_eq!(value, Some("True"));
}
#[test]
fn apply_hub_region_header_omits_when_shared_latch_present_but_false() {
use std::sync::{atomic::AtomicBool, Arc};
let mut request = build_minimal_transport_request();
let shared = Arc::new(AtomicBool::new(false));
let mut state = super::OperationRetryState::initial(0, false, Vec::new(), 3, 1);
state.is_dataplane = true;
state = state.with_shared_hub_region_latch(shared);
super::apply_hub_region_header(&mut request, &state);
let value = request.headers.get_optional_str(&HeaderName::from_static(
request_header_names::HUB_REGION_PROCESSING_ONLY,
));
assert!(value.is_none());
}
fn make_endpoint(region: &'static str) -> CosmosEndpoint {
CosmosEndpoint::regional(
region.into(),
Url::parse(&format!("https://acc-{region}.documents.azure.com:443/"))
.expect("test endpoint url should parse"),
)
}
fn make_advance_test_state(
location_index: usize,
endpoint_count: usize,
) -> crate::driver::pipeline::components::OperationRetryState {
let mut location = LocationIndex::initial(0);
for _ in 0..location_index {
location = location.next(endpoint_count);
}
crate::driver::pipeline::components::OperationRetryState {
location,
failover_retry_count: 0,
session_token_retry_count: 0,
backend_failover_retry_count: 0,
max_failover_retries: 10,
max_backend_failover_retries: 120,
max_session_retries: 2,
can_use_multiple_write_locations: false,
is_dataplane: true,
hub_region_processing_only: false,
shared_hub_region_latch: None,
excluded_regions: Vec::new(),
session_retry_routing:
crate::driver::pipeline::components::SessionRetryRouting::PreferredEndpoints,
partition_key_range_id: None,
ppaf_write_retry_allowed: false,
ppcb_active: false,
pending_write_effects: Vec::new(),
hedge_already_fired: false,
}
}
fn make_advance_test_location(regions: &[&'static str]) -> super::LocationSnapshot {
let endpoints: Vec<CosmosEndpoint> = regions.iter().copied().map(make_endpoint).collect();
let default = endpoints[0].clone();
super::LocationSnapshot::for_tests(Arc::new(AccountEndpointState {
generation: 0,
preferred_read_endpoints: endpoints.clone().into(),
preferred_write_endpoints: endpoints.into(),
unavailable_endpoints: Default::default(),
multiple_write_locations_enabled: false,
default_endpoint: default,
}))
}
fn dummy_last_error() -> crate::error::CosmosError {
crate::error::CosmosError::builder()
.with_message("test-both-transient")
.build()
}
#[test]
fn try_advance_after_both_transient_skips_raced_regions_when_secondary_is_before_primary() {
let regions = ["region-a", "region-b", "region-c", "region-d"];
let location = make_advance_test_location(®ions);
let mut state = make_advance_test_state(2, regions.len());
let primary = crate::options::Region::new("region-c");
let secondary = crate::options::Region::new("region-a");
super::try_advance_after_both_transient(
&mut state,
&location,
true,
Some(&primary),
Some(&secondary),
dummy_last_error(),
)
.expect("budget should allow advance");
let landed = location.account.preferred_read_endpoints[state.location.index()].region();
assert_eq!(
landed.map(crate::options::Region::as_str),
Some("region-d"),
"post-BothTransient LocationIndex must skip the raced primary and \
the raced secondary (no matter where it sat) and land on the \
only untried region",
);
assert_eq!(state.failover_retry_count, 2);
}
#[test]
fn try_advance_after_both_transient_skips_raced_regions_when_secondary_is_after_primary() {
let regions = ["region-a", "region-b", "region-c", "region-d"];
let location = make_advance_test_location(®ions);
let mut state = make_advance_test_state(0, regions.len());
let primary = crate::options::Region::new("region-a");
let secondary = crate::options::Region::new("region-b");
super::try_advance_after_both_transient(
&mut state,
&location,
true,
Some(&primary),
Some(&secondary),
dummy_last_error(),
)
.expect("budget should allow advance");
let landed = location.account.preferred_read_endpoints[state.location.index()].region();
assert_eq!(
landed.map(crate::options::Region::as_str),
Some("region-c"),
"secondary-after-primary case must still skip both raced regions",
);
}
#[test]
fn try_advance_after_both_transient_terminates_when_all_regions_are_raced() {
let regions = ["region-a", "region-b"];
let location = make_advance_test_location(®ions);
let mut state = make_advance_test_state(0, regions.len());
let primary = crate::options::Region::new("region-a");
let secondary = crate::options::Region::new("region-b");
let result = super::try_advance_after_both_transient(
&mut state,
&location,
true,
Some(&primary),
Some(&secondary),
dummy_last_error(),
);
assert!(
result.is_ok(),
"should not surface a terminal error when budget remains"
);
assert_eq!(
state.failover_retry_count, 2,
"two slots are always charged regardless of layout",
);
}
#[test]
fn try_advance_after_both_transient_surfaces_terminal_error_when_budget_exhausted() {
let regions = ["region-a", "region-b", "region-c"];
let location = make_advance_test_location(®ions);
let mut state = make_advance_test_state(0, regions.len());
state.max_failover_retries = 1;
let starting_index = state.location.index();
let primary = crate::options::Region::new("region-a");
let secondary = crate::options::Region::new("region-b");
let result = super::try_advance_after_both_transient(
&mut state,
&location,
true,
Some(&primary),
Some(&secondary),
dummy_last_error(),
);
assert!(result.is_err(), "exhausted budget must surface terminal");
assert_eq!(
state.location.index(),
starting_index,
"exhausted budget must not mutate LocationIndex",
);
assert_eq!(state.failover_retry_count, 0);
}
#[test]
fn both_transient_budget_exhausted_carries_hedge_diagnostics() {
let threshold =
crate::options::HedgeThreshold::new(std::time::Duration::from_millis(200)).unwrap();
let strategy_config = super::HedgingStrategyConfig::new(threshold);
let primary_for_diag = crate::options::Region::new("region-a");
let secondary_for_diag = crate::options::Region::new("region-b");
let race = super::finalize_both_transient(
&crate::models::ActivityId::from_string("both-transient-budget".to_owned()),
None, strategy_config,
Some(crate::options::Region::new("region-a")),
Some(crate::options::Region::new("region-b")),
primary_for_diag.clone(),
secondary_for_diag.clone(),
test_diagnostics(),
None,
false,
);
let super::HedgedRaceResult::BothTransient {
strategy_config: carried_strategy,
primary_region_for_diag,
secondary_region_for_diag,
mut diagnostics,
..
} = race
else {
panic!("deadline-not-elapsed both-transient must return BothTransient");
};
assert_eq!(carried_strategy, strategy_config);
assert_eq!(primary_region_for_diag, primary_for_diag);
assert_eq!(secondary_region_for_diag, secondary_for_diag);
diagnostics.set_hedge_diagnostics(super::HedgeDiagnostics::both_transient(
carried_strategy,
primary_region_for_diag,
secondary_region_for_diag,
false,
));
let ctx = diagnostics.complete();
let hedge = ctx
.hedge_diagnostics()
.expect("budget-exhausted both-transient must carry hedge diagnostics");
assert_eq!(
hedge.terminal_state(),
crate::diagnostics::HedgeTerminalState::BothTransient {
deadline_elapsed: false,
},
);
assert_eq!(hedge.primary_region(), &primary_for_diag);
assert_eq!(hedge.alternate_region(), Some(&secondary_for_diag));
assert_eq!(hedge.response_region(), None);
}
}