use std::collections::{BTreeMap, HashMap, HashSet};
use crate::bridge::envelope::Payload;
use crate::control::gateway::RouteDecision;
use crate::control::server::graph_dispatch::cluster_resolve::resolve_for_vshard;
use crate::control::state::SharedState;
use crate::engine::graph::pattern::executor::{UnresolvedExpansion, VarLenResume, rows_to_msgpack};
use crate::types::{DatabaseId, TenantId, TxnId, VShardId};
use nodedb_cluster::distributed_graph::{
DistributedMatchCoordinator, PatternContinuation, ResolvedContinuationArgs, ShardMatchResult,
};
use super::resume_queue::{PendingResume, resume_seed_key, resume_to_pending};
use super::round_loop::{dispatch_continuations, dispatch_resumes};
use super::round_zero::scatter_round_zero;
const MAX_RESUME_ROUNDS: u32 = 10_000;
const MAX_TRAVERSAL_ROUNDS: usize = 10_000;
pub struct MatchScatterOutcome {
pub rows_payload: Payload,
pub partial: bool,
}
pub(super) struct TaggedShardResult {
pub(super) emitting_node: u64,
pub(super) rows: Vec<HashMap<String, String>>,
pub(super) frontier: Vec<UnresolvedExpansion>,
pub(super) resume: Vec<VarLenResume>,
}
pub async fn scatter_match(
state: &SharedState,
tenant_id: TenantId,
database_id: DatabaseId,
query_bytes: Vec<u8>,
deadline_ms: u64,
txn_id: Option<TxnId>,
) -> crate::Result<MatchScatterOutcome> {
let max_rounds = pattern_round_budget(&query_bytes).max(1) as u32;
let mut coordinator = DistributedMatchCoordinator::new(max_rounds);
let mut partial = false;
let mut pending_resumes: Vec<PendingResume> = Vec::new();
let mut dispatched_seeds: std::collections::HashSet<u64> = std::collections::HashSet::new();
let mut resume_rounds: u32 = 0;
let mut continuations_exhausted = false;
let round0 = scatter_round_zero(
state,
tenant_id,
database_id,
&query_bytes,
deadline_ms,
txn_id,
)
.await?;
for tagged in round0 {
feed_result(
state,
&mut coordinator,
&mut pending_resumes,
&mut dispatched_seeds,
tagged,
)?;
}
while (coordinator.has_pending() && !continuations_exhausted) || !pending_resumes.is_empty() {
if coordinator.has_pending() && !continuations_exhausted {
if !coordinator.advance() {
if coordinator.has_pending() {
partial = true;
}
continuations_exhausted = true;
} else {
let pending = coordinator.take_all_pending();
let tagged = dispatch_continuations(
state,
tenant_id,
database_id,
&query_bytes,
deadline_ms,
txn_id,
pending,
)
.await?;
for t in tagged {
feed_result(
state,
&mut coordinator,
&mut pending_resumes,
&mut dispatched_seeds,
t,
)?;
}
}
}
if !pending_resumes.is_empty() {
if resume_rounds >= MAX_RESUME_ROUNDS {
partial = true;
break;
}
resume_rounds += 1;
let batch = std::mem::take(&mut pending_resumes);
let tagged = dispatch_resumes(
state,
tenant_id,
database_id,
&query_bytes,
deadline_ms,
txn_id,
batch,
)
.await?;
for t in tagged {
feed_result(
state,
&mut coordinator,
&mut pending_resumes,
&mut dispatched_seeds,
t,
)?;
}
}
}
let rows_payload = dedup_and_encode(&coordinator.completed)?;
Ok(MatchScatterOutcome {
rows_payload,
partial,
})
}
pub(super) fn feed_result(
state: &SharedState,
coordinator: &mut DistributedMatchCoordinator,
pending_resumes: &mut Vec<PendingResume>,
dispatched_seeds: &mut std::collections::HashSet<u64>,
tagged: TaggedShardResult,
) -> crate::Result<()> {
let TaggedShardResult {
emitting_node,
rows,
frontier,
resume,
} = tagged;
for cursor in resume {
if !dispatched_seeds.insert(resume_seed_key(&cursor)) {
continue;
}
if let Some(pending) = resume_to_pending(state, cursor)? {
pending_resumes.push(pending);
}
}
let continuations = frontier_to_continuations(state, emitting_node, frontier)?;
coordinator.add_shard_result(ShardMatchResult {
shard_id: emitting_node as u32,
completed_rows: rows,
continuations,
});
Ok(())
}
fn frontier_to_continuations(
state: &SharedState,
emitting_node: u64,
frontier: Vec<UnresolvedExpansion>,
) -> crate::Result<Vec<PatternContinuation>> {
let mut out = Vec::new();
for entry in frontier {
let target_vshard = VShardId::from_key(entry.node_name.as_bytes()).as_u32();
let decision = resolve_for_vshard(state, target_vshard);
let owner_node = match decision {
RouteDecision::Local => state.node_id,
RouteDecision::Remote { node_id, .. } => node_id,
RouteDecision::LeaderUnknown { vshard_id } => {
return Err(crate::Error::NotLeader {
vshard_id: VShardId::new((vshard_id % VShardId::COUNT as u64) as u32),
leader_node: 0,
leader_addr: String::new(),
});
}
RouteDecision::Broadcast { .. } => {
return Err(crate::Error::Internal {
detail: "match scatter: resolve_decision returned Broadcast for a \
single vShard"
.into(),
});
}
};
if owner_node == emitting_node {
continue;
}
out.push(PatternContinuation::from_resolved(
ResolvedContinuationArgs {
target_shard: target_vshard,
source_shard: emitting_node as u32,
bindings: entry.partial_row,
next_triple_idx: entry.triple_idx,
start_node: entry.node_name,
start_binding: entry.binding_var,
},
));
}
Ok(out)
}
pub(super) fn decode_rows(payload: &Payload) -> crate::Result<Vec<HashMap<String, String>>> {
if payload.is_empty() {
return Ok(Vec::new());
}
zerompk::from_msgpack::<Vec<HashMap<String, String>>>(payload.as_ref()).map_err(|e| {
crate::Error::Codec {
detail: format!("match scatter: invalid rows array: {e}"),
}
})
}
fn dedup_and_encode(rows: &[HashMap<String, String>]) -> crate::Result<Payload> {
let mut seen: HashSet<Vec<(String, String)>> = HashSet::new();
let mut deduped: Vec<HashMap<String, String>> = Vec::with_capacity(rows.len());
for row in rows {
let fingerprint: Vec<(String, String)> = row
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect::<BTreeMap<_, _>>()
.into_iter()
.collect();
if seen.insert(fingerprint) {
deduped.push(row.clone());
}
}
let bytes = rows_to_msgpack(&deduped)?;
Ok(Payload::from_vec(bytes))
}
fn pattern_round_budget(query_bytes: &[u8]) -> usize {
use crate::engine::graph::pattern::ast::MatchQuery;
let query: MatchQuery = match zerompk::from_msgpack(query_bytes) {
Ok(q) => q,
Err(_) => return 0,
};
query
.clauses
.iter()
.flat_map(|c| c.patterns.iter())
.flat_map(|chain| chain.triples.iter())
.map(|triple| triple.edge.max_hops.max(1))
.fold(0usize, |acc, hops| acc.saturating_add(hops))
.min(MAX_TRAVERSAL_ROUNDS)
}