Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
219 changes: 193 additions & 26 deletions include/sql_engine/distributed_planner.h
Original file line number Diff line number Diff line change
Expand Up @@ -206,6 +206,13 @@ class DistributedPlanner {
return result;
}

case PlanNodeType::DERIVED_SCAN: {
PlanNode* result = make_plan_node(arena_, PlanNodeType::DERIVED_SCAN);
result->derived_scan = node->derived_scan;
result->derived_scan.inner_plan = distribute(node->derived_scan.inner_plan);
return result;
}

case PlanNodeType::SET_OP: {
PlanNode* result = make_plan_node(arena_, PlanNodeType::SET_OP);
result->set_op = node->set_op;
Expand Down Expand Up @@ -293,6 +300,9 @@ class DistributedPlanner {
if (contains_type(node->merge_sort.children[i], type)) return true;
}
}
if (node->type == PlanNodeType::DERIVED_SCAN) {
return contains_type(node->derived_scan.inner_plan, type);
}
return false;
}

Expand Down Expand Up @@ -545,7 +555,7 @@ class DistributedPlanner {
std::memcpy(bn, backend, blen + 1);
node->remote_scan.backend_name = bn;
node->remote_scan.remote_sql = sql.ptr;
node->remote_scan.remote_sql_len = static_cast<uint16_t>(sql.len);
node->remote_scan.remote_sql_len = sql.len;
node->remote_scan.table = table;
// Caller is responsible for setting output_exprs when the remote SQL
// is not a passthrough SELECT *. make_plan_node() already zero-fills
Expand Down Expand Up @@ -782,15 +792,53 @@ class DistributedPlanner {
}

// Case 4: Distributed sort + limit
PlanNode* local_sort(PlanNode* sort_node) {
PlanNode* result = make_plan_node(arena_, PlanNodeType::SORT);
result->sort = sort_node->sort;
result->left = distribute_node(sort_node->left);
return result;
}

int sort_key_table_ordinal(const sql_parser::AstNode* key, const TableInfo* table) const {
if (!key || !table) return -1;
if (key->type == sql_parser::NodeType::NODE_LITERAL_INT) {
sql_parser::StringRef sv = key->value();
if (!sv.ptr || sv.len == 0) return -1;
int64_t n = std::strtoll(sv.ptr, nullptr, 10);
if (n < 1 || n > static_cast<int64_t>(table->column_count)) return -1;
return static_cast<int>(n - 1);
}
sql_parser::StringRef col_name;
if (key->type == sql_parser::NodeType::NODE_COLUMN_REF ||
key->type == sql_parser::NodeType::NODE_IDENTIFIER) {
col_name = key->value();
} else if (key->type == sql_parser::NodeType::NODE_QUALIFIED_NAME) {
const sql_parser::AstNode* c = key->first_child;
if (c && c->next_sibling) col_name = c->next_sibling->value();
else if (c) col_name = c->value();
} else {
return -1;
}
if (!col_name.ptr) return -1;
const ColumnInfo* col = catalog_.get_column(table, col_name);
if (!col) return -1;
return static_cast<int>(col->ordinal);
}

bool all_sort_keys_are_table_columns(const PlanNode* sort_node, const TableInfo* table) const {
if (!sort_node || !table) return false;
for (uint16_t i = 0; i < sort_node->sort.count; ++i) {
if (sort_key_table_ordinal(sort_node->sort.keys[i], table) < 0) return false;
}
return true;
}

PlanNode* distribute_sort(PlanNode* sort_node) {
if (contains_type(sort_node->left, PlanNodeType::WINDOW) ||
contains_type(sort_node->left, PlanNodeType::DERIVED_SCAN) ||
contains_type(sort_node->left, PlanNodeType::AGGREGATE) ||
contains_type(sort_node->left, PlanNodeType::MERGE_AGGREGATE)) {
PlanNode* result = make_plan_node(arena_, PlanNodeType::SORT);
result->sort = sort_node->sort;
result->left = distribute_node(sort_node->left);
return result;
return local_sort(sort_node);
}

ScanContext ctx = extract_scan_context(sort_node->left);
Expand All @@ -809,6 +857,10 @@ class DistributedPlanner {
return result;
}

if (!all_sort_keys_are_table_columns(sort_node, table)) {
return local_sort(sort_node);
}

if (!shards_.is_sharded(table->table_name)) {
// Unsharded -- push sort to remote
sql_parser::StringRef sql = qb_.build_select(
Expand Down Expand Up @@ -875,8 +927,12 @@ class DistributedPlanner {
const TableInfo* table = ctx.scan->scan.table;
if (shards_.has_table(table->table_name) &&
shards_.is_sharded(table->table_name)) {
// Case 4: Sharded sort + limit
// Each shard: ORDER BY + LIMIT, MergeSort, then outer Limit
if (!all_sort_keys_are_table_columns(sort_node, table)) {
PlanNode* result = make_plan_node(arena_, PlanNodeType::LIMIT);
result->limit = limit_node->limit;
result->left = distribute_node(limit_node->left);
return result;
}
int64_t remote_limit = limit_node->limit.count + limit_node->limit.offset;

PlanNode* merge = make_sharded_merge_sort(
Expand Down Expand Up @@ -943,12 +999,72 @@ class DistributedPlanner {
return result;
}

// Case 5: Cross-backend join
bool join_on_shard_keys(const sql_parser::AstNode* cond,
sql_parser::StringRef left_key,
sql_parser::StringRef right_key) const {
if (!cond || cond->type != sql_parser::NodeType::NODE_BINARY_OP) return false;
sql_parser::StringRef op = cond->value();
if (op.len != 1 || op.ptr[0] != '=') return false;
const sql_parser::AstNode* l = cond->first_child;
const sql_parser::AstNode* r = l ? l->next_sibling : nullptr;
if (!l || !r) return false;
return (is_shard_key_ref(l, left_key) && is_shard_key_ref(r, right_key)) ||
(is_shard_key_ref(l, right_key) && is_shard_key_ref(r, left_key));
}

PlanNode* distribute_colocated_join(PlanNode* join_node,
const TableInfo* left_table,
const TableInfo* right_table) {
ScanContext lctx = extract_scan_context(join_node->left);
ScanContext rctx = extract_scan_context(join_node->right);
const sql_parser::AstNode* where_expr = nullptr;
if (lctx.where_expr && rctx.where_expr) {
sql_parser::AstNode* and_node = sql_parser::make_node(
arena_, sql_parser::NodeType::NODE_BINARY_OP,
sql_parser::StringRef{"AND", 3});
and_node->add_child(const_cast<sql_parser::AstNode*>(lctx.where_expr));
and_node->add_child(const_cast<sql_parser::AstNode*>(rctx.where_expr));
where_expr = and_node;
} else if (lctx.where_expr) {
where_expr = lctx.where_expr;
} else {
where_expr = rctx.where_expr;
}

const auto& shard_list = shards_.get_shards(left_table->table_name);
PlanNode* current = nullptr;
for (const auto& shard : shard_list) {
sql_parser::StringRef sql = qb_.build_select_join(
left_table, right_table, join_node->join.condition, where_expr);
PlanNode* rs = make_remote_scan(shard.backend_name.c_str(), sql, left_table);
if (!current) {
current = rs;
} else {
PlanNode* union_node = make_plan_node(arena_, PlanNodeType::SET_OP);
union_node->set_op.op = SET_OP_UNION;
union_node->set_op.all = true;
union_node->left = current;
union_node->right = rs;
current = union_node;
}
}
return current ? current : join_node;
}

PlanNode* distribute_join(PlanNode* join_node) {
// Get tables from each side
const TableInfo* left_table = find_table(join_node->left);
const TableInfo* right_table = find_table(join_node->right);

if (left_table && right_table &&
shards_.is_sharded(left_table->table_name) &&
shards_.is_sharded(right_table->table_name) &&
shards_.same_routing(left_table->table_name, right_table->table_name) &&
join_on_shard_keys(join_node->join.condition,
shards_.get_shard_key(left_table->table_name),
shards_.get_shard_key(right_table->table_name))) {
return distribute_colocated_join(join_node, left_table, right_table);
}

PlanNode* left_dist = nullptr;
PlanNode* right_dist = nullptr;

Expand Down Expand Up @@ -1116,18 +1232,13 @@ class DistributedPlanner {
PlanNode* distribute_update(PlanNode* plan) {
const auto& up = plan->update_plan;
const TableInfo* table = up.table;
if (!table || !shards_.has_table(table->table_name)) return plan;

// Multi-table UPDATE: emit full SQL from AST, route to primary table's backend
if (up.original_ast) {
sql_parser::StringRef sql = qb_.build_update_from_ast(up.original_ast);
if (!shards_.is_sharded(table->table_name)) {
return make_remote_scan(shards_.get_backend(table->table_name), sql, table);
}
const auto& shard_list = shards_.get_shards(table->table_name);
return scatter_dml_to_shards(table, shard_list, [&]() { return sql; });
return distribute_multi_table_dml(up.original_ast, table, true);
}

if (!table || !shards_.has_table(table->table_name)) return plan;

// Check for cross-shard subqueries in WHERE and rewrite
const sql_parser::AstNode* where_expr = up.where_expr;
if (where_expr && has_subquery(where_expr) && remote_executor_) {
Expand Down Expand Up @@ -1164,18 +1275,13 @@ class DistributedPlanner {
PlanNode* distribute_delete(PlanNode* plan) {
const auto& dp = plan->delete_plan;
const TableInfo* table = dp.table;
if (!table || !shards_.has_table(table->table_name)) return plan;

// Multi-table DELETE: emit full SQL from AST, route to primary table's backend
if (dp.original_ast) {
sql_parser::StringRef sql = qb_.build_delete_from_ast(dp.original_ast);
if (!shards_.is_sharded(table->table_name)) {
return make_remote_scan(shards_.get_backend(table->table_name), sql, table);
}
const auto& shard_list = shards_.get_shards(table->table_name);
return scatter_dml_to_shards(table, shard_list, [&]() { return sql; });
return distribute_multi_table_dml(dp.original_ast, table, false);
}

if (!table || !shards_.has_table(table->table_name)) return plan;

// Check for cross-shard subqueries in WHERE and rewrite
const sql_parser::AstNode* where_expr = dp.where_expr;
if (where_expr && has_subquery(where_expr) && remote_executor_) {
Expand Down Expand Up @@ -1262,7 +1368,68 @@ class DistributedPlanner {
return false;
}

// Scatter DML SQL to all shards, combining results via UNION ALL
void collect_ast_table_names(const sql_parser::AstNode* n,
std::vector<sql_parser::StringRef>& out) const {
if (!n) return;
if (n->type == sql_parser::NodeType::NODE_TABLE_REF && n->first_child) {
const sql_parser::AstNode* name = n->first_child;
if (name->type == sql_parser::NodeType::NODE_IDENTIFIER) {
out.push_back(name->value());
} else if (name->type == sql_parser::NodeType::NODE_QUALIFIED_NAME) {
const sql_parser::AstNode* schema = name->first_child;
const sql_parser::AstNode* table = schema ? schema->next_sibling : nullptr;
if (table) out.push_back(table->value());
else if (schema) out.push_back(schema->value());
}
}
for (const sql_parser::AstNode* c = n->first_child; c; c = c->next_sibling) {
collect_ast_table_names(c, out);
}
}

PlanNode* distribute_multi_table_dml(const sql_parser::AstNode* ast,
const TableInfo* primary,
bool is_update) {
std::vector<sql_parser::StringRef> names;
collect_ast_table_names(ast, names);
const char* backend = nullptr;
bool saw_mapped = false;
for (sql_parser::StringRef name : names) {
if (!shards_.has_table(name)) continue;
saw_mapped = true;
if (shards_.is_sharded(name)) {
return fail_dml(is_update
? "multi-table UPDATE is not supported on sharded tables"
: "multi-table DELETE is not supported on sharded tables");
}
const char* b = shards_.get_backend(name);
if (backend && b && std::strcmp(backend, b) != 0) {
return fail_dml(is_update
? "multi-table UPDATE spans multiple backends"
: "multi-table DELETE spans multiple backends");
}
if (b) backend = b;
}
if (!backend && primary && shards_.has_table(primary->table_name)) {
if (shards_.is_sharded(primary->table_name)) {
return fail_dml(is_update
? "multi-table UPDATE is not supported on sharded tables"
: "multi-table DELETE is not supported on sharded tables");
}
backend = shards_.get_backend(primary->table_name);
saw_mapped = true;
}
if (!backend || !saw_mapped) {
return fail_dml(is_update
? "multi-table UPDATE is not supported on sharded tables"
: "multi-table DELETE is not supported on sharded tables");
}
sql_parser::StringRef sql = is_update
? qb_.build_update_from_ast(ast)
: qb_.build_delete_from_ast(ast);
return make_remote_scan(backend, sql, primary);
}

PlanNode* scatter_dml_to_shards(const TableInfo* table,
const std::vector<ShardInfo>& shard_list,
std::function<sql_parser::StringRef()> build_sql) {
Expand Down
12 changes: 12 additions & 0 deletions include/sql_engine/distributed_txn.h
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,18 @@ class DistributedTransactionManager : public TransactionManager {
return executor_.execute_dml(backend_name, sql);
}

ResultSet route_query(const char* backend_name,
sql_parser::StringRef sql) override {
if (!active_) return executor_.execute(backend_name, sql);
auto it = sessions_.find(backend_name);
if (it != sessions_.end() && it->second) {
return it->second->execute(sql);
}
return executor_.execute(backend_name, sql);
}

bool route_query_supported() const override { return true; }

bool commit() override {
if (!active_) return false;
if (participants_.empty()) {
Expand Down
Loading