use serde_json::{Map, Value as JsonValue};
use crate::bridge::envelope::PhysicalPlan;
use crate::control::security::identity::AuthenticatedIdentity;
use crate::control::server::response_shape::types::ShapedRows;
use crate::control::state::SharedState;
use crate::engine::graph::traversal_options::GraphTraversalOptions;
use crate::types::{DatabaseId, TraceId, VShardId};
use super::super::super::result::{DdlError, DdlResult};
use super::parse::{extract_function_args, extract_number_after, json_to_decimal};
use super::support::ddl_err;
pub async fn tree_sum(
state: &SharedState,
identity: &AuthenticatedIdentity,
database_id: DatabaseId,
sql: &str,
) -> Result<Vec<DdlResult>, DdlError> {
let tenant_id = identity.tenant_id;
let upper = sql.to_uppercase();
let args = extract_function_args(&upper, sql, "TREE_SUM")?;
if args.len() < 3 {
return Err(ddl_err(
"42601",
"TREE_SUM requires (column, graph_index, root_id [, collection])",
));
}
let sum_column = args[0].trim().to_lowercase();
let graph_index = args[1].trim().to_lowercase();
let root_id = args[2]
.trim()
.trim_matches('\'')
.trim_matches('"')
.to_string();
let explicit_collection = args
.get(3)
.map(|s| s.trim().trim_matches('\'').trim_matches('"').to_lowercase());
let max_depth = extract_number_after(&upper, "MAX_DEPTH")?.unwrap_or(100);
let dir = crate::engine::graph::edge_store::Direction::Out;
let bfs_result = crate::control::server::graph_dispatch::cross_core_bfs_with_options(
state,
crate::control::server::graph_dispatch::CrossCoreBfsParams {
tenant_id,
database_id,
start_nodes: vec![root_id.clone()],
edge_label: Some(graph_index),
direction: dir,
max_depth,
options: &GraphTraversalOptions::default(),
},
)
.await
.map_err(|e| ddl_err("XX000", format!("BFS failed: {e}")))?;
let bfs_json =
crate::data::executor::response_codec::decode_payload_to_json(&bfs_result.payload);
let bfs_nodes: Vec<String> = sonic_rs::from_str::<Vec<serde_json::Value>>(&bfs_json)
.unwrap_or_default()
.into_iter()
.filter_map(|v| v.as_str().map(String::from))
.collect();
let mut all_ids: Vec<String> = vec![root_id];
for id in bfs_nodes {
if !id.is_empty() && !all_ids.contains(&id) {
all_ids.push(id);
}
}
let mut total = rust_decimal::Decimal::ZERO;
let collections_to_search: Vec<String> = if let Some(ref coll) = explicit_collection {
vec![coll.clone()]
} else {
state
.credentials
.catalog()
.load_collections_for_tenant(database_id, tenant_id.as_u64())
.unwrap_or_default()
.iter()
.map(|c| c.name.clone())
.collect()
};
for node_id in &all_ids {
for coll_name in &collections_to_search {
let coll_vshard = VShardId::from_collection_in_database(database_id, coll_name);
let pk_bytes = node_id.as_bytes().to_vec();
let surrogate = state
.surrogate_assigner
.lookup(database_id, tenant_id, coll_name, &pk_bytes)
.map_err(|e| ddl_err("XX000", format!("surrogate lookup: {e}")))?
.unwrap_or(nodedb_types::Surrogate::ZERO);
let get_plan =
PhysicalPlan::Document(nodedb_physical::physical_plan::DocumentOp::PointGet {
collection: coll_name.clone(),
document_id: node_id.clone(),
surrogate,
pk_bytes,
rls_filters: Vec::new(),
system_time: nodedb_types::SystemTimeScope::Current,
valid_at_ms: None,
});
if let Ok(resp) = crate::control::server::dispatch_utils::dispatch_to_data_plane(
state,
tenant_id,
database_id,
coll_vshard,
get_plan,
TraceId::ZERO,
)
.await
{
let doc_json =
crate::data::executor::response_codec::decode_payload_to_json(&resp.payload);
if let Ok(doc) = sonic_rs::from_str::<serde_json::Value>(&doc_json)
&& let Some(val) = doc.get(&sum_column)
{
total += json_to_decimal(val);
break; }
}
}
}
let mut row = Map::new();
row.insert("tree_sum".to_string(), JsonValue::String(total.to_string()));
Ok(vec![DdlResult::Rows(ShapedRows {
columns: vec!["tree_sum".to_string()],
column_types: ShapedRows::text_types(1),
rows: vec![row],
notice: None,
})])
}