use futures::future::join_all;
use std::time::{Duration, Instant};
use crate::bridge::envelope::{PhysicalPlan, Priority, Request, Response, Status};
use crate::control::arrow_convert;
use crate::control::gateway::core::QueryContext;
use crate::control::server::payload_merge::{encode_msgpack_array, extract_msgpack_elements};
use crate::control::server::result_stream::ResultStream;
use crate::control::state::SharedState;
use crate::types::{
DatabaseId, Lsn, ReadConsistency, RequestId, TenantId, TraceId, TxnId, VShardId,
};
pub(crate) fn eager_dispatch_to_all_cores(
state: &SharedState,
tenant_id: TenantId,
database_id: DatabaseId,
trace_id: TraceId,
txn_id: Option<TxnId>,
plan_for_core: impl Fn(usize) -> PhysicalPlan,
) -> crate::Result<
Vec<(
usize,
tokio::sync::mpsc::Receiver<crate::bridge::envelope::Response>,
)>,
> {
let deadline_secs = state.tuning.network.default_deadline_secs;
let num_cores = state
.dispatcher
.lock()
.unwrap_or_else(|p| p.into_inner())
.num_cores();
let mut receivers = Vec::with_capacity(num_cores);
for core_id in 0..num_cores {
let request_id = state.next_request_id();
let vshard_id = VShardId::new(core_id as u32);
let request = Request {
request_id,
tenant_id,
database_id,
vshard_id,
plan: plan_for_core(core_id),
deadline: Instant::now() + Duration::from_secs(deadline_secs),
priority: Priority::Normal,
trace_id,
consistency: ReadConsistency::Strong,
idempotency_key: None,
event_source: crate::event::EventSource::User,
user_roles: Vec::new(),
user_id: None,
statement_digest: None,
txn_id,
wal_lsn: None,
resolved_now_ms: None,
admission: crate::bridge::envelope::Admission::Exempt(
crate::bridge::envelope::ExemptReason::Read,
),
};
let rx = state.tracker.register(request_id);
state
.dispatcher
.lock()
.unwrap_or_else(|p| p.into_inner())
.dispatch_to_core(core_id, request)?;
receivers.push((core_id, rx));
}
Ok(receivers)
}
pub struct GatherOutcome {
pub raw: Vec<u8>,
pub merged_array: Vec<u8>,
pub watermark_lsn: Lsn,
pub read_version_lsn: Lsn,
pub shard_watermarks: Vec<(VShardId, Lsn)>,
}
pub async fn gather_all_cores(
state: &SharedState,
tenant_id: TenantId,
database_id: DatabaseId,
plan: PhysicalPlan,
trace_id: TraceId,
txn_id: Option<TxnId>,
) -> crate::Result<GatherOutcome> {
crate::control::server::broadcast::broadcast_call_count_increment();
let deadline_secs = state.tuning.network.default_deadline_secs;
let receivers =
eager_dispatch_to_all_cores(state, tenant_id, database_id, trace_id, txn_id, |_| {
plan.clone()
})?;
let deadline = Duration::from_secs(deadline_secs);
let max_result_bytes = state.tuning.network.max_query_result_bytes as usize;
let response_futures = receivers.into_iter().map(|(core_id, mut rx)| async move {
let result = match tokio::time::timeout(
deadline,
crate::control::server::dispatch_utils::collect_bounded_response(
&mut rx,
max_result_bytes,
),
)
.await
{
Err(_) => Err(crate::Error::Dispatch {
detail: format!("gather timeout on core {core_id}"),
}),
Ok(Ok(resp)) => Ok(resp),
Ok(Err(crate::control::server::dispatch_utils::DispatchCollectError::OverBudget {
bytes,
})) => Err(crate::Error::ExecutionLimitExceeded {
detail: format!(
"gather on core {core_id} exceeded max_query_result_bytes \
({bytes} > {max_result_bytes} bytes)"
),
}),
Ok(Err(
crate::control::server::dispatch_utils::DispatchCollectError::ChannelClosed,
)) => Err(crate::Error::Dispatch {
detail: format!("gather channel closed on core {core_id}"),
}),
};
(core_id, result)
});
let results: Vec<(usize, crate::Result<Response>)> = join_all(response_futures).await;
let mut raw = Vec::new();
let mut all_elements: Vec<Vec<u8>> = Vec::new();
let mut max_lsn = Lsn::ZERO;
let mut max_read_version = Lsn::ZERO;
let mut shard_watermarks: Vec<(VShardId, Lsn)> = Vec::new();
let mut had_error = false;
let mut error_msg = String::new();
for (core_id, result) in results {
let resp = match result {
Ok(r) => r,
Err(e) => {
had_error = true;
error_msg = e.to_string();
continue;
}
};
if resp.status == Status::Error {
if let Some(ec) = resp.error_code.as_deref() {
match ec {
crate::bridge::envelope::ErrorCode::NotFound => continue,
_ => {
had_error = true;
error_msg = format!("{ec:?}");
}
}
}
continue;
}
shard_watermarks.push((VShardId::new(core_id as u32), resp.watermark_lsn));
if resp.watermark_lsn > max_lsn {
max_lsn = resp.watermark_lsn;
}
if resp.read_version_lsn > max_read_version {
max_read_version = resp.read_version_lsn;
}
if resp.payload.is_empty() {
continue;
}
let payload_bytes: &[u8] = resp.payload.as_ref();
raw.extend_from_slice(payload_bytes);
all_elements.extend(extract_msgpack_elements(payload_bytes));
}
if had_error && all_elements.is_empty() && raw.is_empty() {
return Err(crate::Error::Dispatch { detail: error_msg });
}
let merged_array = encode_msgpack_array(&all_elements);
Ok(GatherOutcome {
raw,
merged_array,
watermark_lsn: max_lsn,
read_version_lsn: max_read_version,
shard_watermarks,
})
}
pub fn gather_all_cores_stream(
state: &SharedState,
tenant_id: TenantId,
database_id: DatabaseId,
plan: PhysicalPlan,
trace_id: TraceId,
txn_id: Option<TxnId>,
) -> crate::Result<ResultStream> {
use crate::control::server::result_stream::stream_response_channel;
crate::control::server::broadcast::broadcast_call_count_increment();
let max_result_bytes = state.tuning.network.max_query_result_bytes as usize;
let per_core: Vec<ResultStream> =
eager_dispatch_to_all_cores(state, tenant_id, database_id, trace_id, txn_id, |_| {
plan.clone()
})?
.into_iter()
.map(|(_core_id, rx)| stream_response_channel(rx, max_result_bytes, true))
.collect();
Ok(Box::pin(futures::stream::select_all(per_core)))
}
pub async fn gather_all_vshards(
state: &SharedState,
tenant_id: TenantId,
database_id: DatabaseId,
plan: PhysicalPlan,
trace_id: TraceId,
txn_id: Option<TxnId>,
) -> crate::Result<GatherOutcome> {
let Some(gateway) = state.gateway.get() else {
return super::owning_core::gather_single_node(
state,
tenant_id,
database_id,
plan,
trace_id,
txn_id,
)
.await;
};
if nodedb_physical::physical_plan::plan_contains_cluster_partitioned_leaf(&plan) {
return gather_all_cores(state, tenant_id, database_id, plan, trace_id, txn_id).await;
}
let ctx = QueryContext {
tenant_id,
trace_id,
database_id,
txn_id,
};
let (payloads, shard_watermarks, read_version_lsn): (Vec<Vec<u8>>, Vec<(VShardId, Lsn)>, Lsn) =
Box::pin(gateway.execute_with_watermarks(&ctx, plan))
.await
.map_err(|e| crate::Error::Dispatch {
detail: format!("cross-node gather via gateway: {e}"),
})?;
let mut all_elements: Vec<Vec<u8>> = Vec::new();
let mut raw = Vec::new();
for payload in &payloads {
raw.extend_from_slice(payload);
all_elements.extend(extract_msgpack_elements(payload));
}
let merged_array = encode_msgpack_array(&all_elements);
let watermark_lsn = shard_watermarks
.iter()
.map(|(_, lsn)| *lsn)
.max()
.unwrap_or(Lsn::ZERO);
Ok(GatherOutcome {
raw,
merged_array,
watermark_lsn,
read_version_lsn,
shard_watermarks,
})
}
pub fn finalize_aggregate(merged_array: &[u8]) -> Vec<u8> {
if let Some(batch) = arrow_convert::msgpack_rows_to_record_batch(merged_array) {
tracing::trace!(
rows = batch.num_rows(),
columns = batch.num_columns(),
"arrow aggregate post-processing: merged {} rows",
batch.num_rows(),
);
}
merged_array.to_vec()
}
pub(crate) async fn stream_to_response(stream: ResultStream) -> crate::Result<Response> {
let (merged, watermark_lsn) =
crate::control::server::result_stream::materialize(stream).await?;
Ok(outcome_to_response(merged, watermark_lsn, Lsn::ZERO))
}
pub(super) fn outcome_to_response(
merged_array: Vec<u8>,
watermark_lsn: Lsn,
read_version_lsn: Lsn,
) -> Response {
Response {
request_id: RequestId::new(0),
status: Status::Ok,
attempt: 1,
partial: false,
payload: crate::bridge::envelope::Payload::from_vec(merged_array),
watermark_lsn,
error_code: None,
read_set_valid: None,
read_version_lsn,
write_set: Vec::new(),
}
}