Skip to content

Commit ec89e65

Browse files
authored
Merge pull request #58 from ProxySQL/fix/shard-derived-sort-sql
fix: distribute derived tables and stop silent sort/SQL bugs
2 parents 5da29e2 + f1e896f commit ec89e65

14 files changed

Lines changed: 684 additions & 40 deletions

include/sql_engine/distributed_planner.h

Lines changed: 193 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -206,6 +206,13 @@ class DistributedPlanner {
206206
return result;
207207
}
208208

209+
case PlanNodeType::DERIVED_SCAN: {
210+
PlanNode* result = make_plan_node(arena_, PlanNodeType::DERIVED_SCAN);
211+
result->derived_scan = node->derived_scan;
212+
result->derived_scan.inner_plan = distribute(node->derived_scan.inner_plan);
213+
return result;
214+
}
215+
209216
case PlanNodeType::SET_OP: {
210217
PlanNode* result = make_plan_node(arena_, PlanNodeType::SET_OP);
211218
result->set_op = node->set_op;
@@ -293,6 +300,9 @@ class DistributedPlanner {
293300
if (contains_type(node->merge_sort.children[i], type)) return true;
294301
}
295302
}
303+
if (node->type == PlanNodeType::DERIVED_SCAN) {
304+
return contains_type(node->derived_scan.inner_plan, type);
305+
}
296306
return false;
297307
}
298308

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

784794
// Case 4: Distributed sort + limit
795+
PlanNode* local_sort(PlanNode* sort_node) {
796+
PlanNode* result = make_plan_node(arena_, PlanNodeType::SORT);
797+
result->sort = sort_node->sort;
798+
result->left = distribute_node(sort_node->left);
799+
return result;
800+
}
801+
802+
int sort_key_table_ordinal(const sql_parser::AstNode* key, const TableInfo* table) const {
803+
if (!key || !table) return -1;
804+
if (key->type == sql_parser::NodeType::NODE_LITERAL_INT) {
805+
sql_parser::StringRef sv = key->value();
806+
if (!sv.ptr || sv.len == 0) return -1;
807+
int64_t n = std::strtoll(sv.ptr, nullptr, 10);
808+
if (n < 1 || n > static_cast<int64_t>(table->column_count)) return -1;
809+
return static_cast<int>(n - 1);
810+
}
811+
sql_parser::StringRef col_name;
812+
if (key->type == sql_parser::NodeType::NODE_COLUMN_REF ||
813+
key->type == sql_parser::NodeType::NODE_IDENTIFIER) {
814+
col_name = key->value();
815+
} else if (key->type == sql_parser::NodeType::NODE_QUALIFIED_NAME) {
816+
const sql_parser::AstNode* c = key->first_child;
817+
if (c && c->next_sibling) col_name = c->next_sibling->value();
818+
else if (c) col_name = c->value();
819+
} else {
820+
return -1;
821+
}
822+
if (!col_name.ptr) return -1;
823+
const ColumnInfo* col = catalog_.get_column(table, col_name);
824+
if (!col) return -1;
825+
return static_cast<int>(col->ordinal);
826+
}
827+
828+
bool all_sort_keys_are_table_columns(const PlanNode* sort_node, const TableInfo* table) const {
829+
if (!sort_node || !table) return false;
830+
for (uint16_t i = 0; i < sort_node->sort.count; ++i) {
831+
if (sort_key_table_ordinal(sort_node->sort.keys[i], table) < 0) return false;
832+
}
833+
return true;
834+
}
835+
785836
PlanNode* distribute_sort(PlanNode* sort_node) {
786837
if (contains_type(sort_node->left, PlanNodeType::WINDOW) ||
787838
contains_type(sort_node->left, PlanNodeType::DERIVED_SCAN) ||
788839
contains_type(sort_node->left, PlanNodeType::AGGREGATE) ||
789840
contains_type(sort_node->left, PlanNodeType::MERGE_AGGREGATE)) {
790-
PlanNode* result = make_plan_node(arena_, PlanNodeType::SORT);
791-
result->sort = sort_node->sort;
792-
result->left = distribute_node(sort_node->left);
793-
return result;
841+
return local_sort(sort_node);
794842
}
795843

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

860+
if (!all_sort_keys_are_table_columns(sort_node, table)) {
861+
return local_sort(sort_node);
862+
}
863+
812864
if (!shards_.is_sharded(table->table_name)) {
813865
// Unsharded -- push sort to remote
814866
sql_parser::StringRef sql = qb_.build_select(
@@ -875,8 +927,12 @@ class DistributedPlanner {
875927
const TableInfo* table = ctx.scan->scan.table;
876928
if (shards_.has_table(table->table_name) &&
877929
shards_.is_sharded(table->table_name)) {
878-
// Case 4: Sharded sort + limit
879-
// Each shard: ORDER BY + LIMIT, MergeSort, then outer Limit
930+
if (!all_sort_keys_are_table_columns(sort_node, table)) {
931+
PlanNode* result = make_plan_node(arena_, PlanNodeType::LIMIT);
932+
result->limit = limit_node->limit;
933+
result->left = distribute_node(limit_node->left);
934+
return result;
935+
}
880936
int64_t remote_limit = limit_node->limit.count + limit_node->limit.offset;
881937

882938
PlanNode* merge = make_sharded_merge_sort(
@@ -943,12 +999,72 @@ class DistributedPlanner {
943999
return result;
9441000
}
9451001

946-
// Case 5: Cross-backend join
1002+
bool join_on_shard_keys(const sql_parser::AstNode* cond,
1003+
sql_parser::StringRef left_key,
1004+
sql_parser::StringRef right_key) const {
1005+
if (!cond || cond->type != sql_parser::NodeType::NODE_BINARY_OP) return false;
1006+
sql_parser::StringRef op = cond->value();
1007+
if (op.len != 1 || op.ptr[0] != '=') return false;
1008+
const sql_parser::AstNode* l = cond->first_child;
1009+
const sql_parser::AstNode* r = l ? l->next_sibling : nullptr;
1010+
if (!l || !r) return false;
1011+
return (is_shard_key_ref(l, left_key) && is_shard_key_ref(r, right_key)) ||
1012+
(is_shard_key_ref(l, right_key) && is_shard_key_ref(r, left_key));
1013+
}
1014+
1015+
PlanNode* distribute_colocated_join(PlanNode* join_node,
1016+
const TableInfo* left_table,
1017+
const TableInfo* right_table) {
1018+
ScanContext lctx = extract_scan_context(join_node->left);
1019+
ScanContext rctx = extract_scan_context(join_node->right);
1020+
const sql_parser::AstNode* where_expr = nullptr;
1021+
if (lctx.where_expr && rctx.where_expr) {
1022+
sql_parser::AstNode* and_node = sql_parser::make_node(
1023+
arena_, sql_parser::NodeType::NODE_BINARY_OP,
1024+
sql_parser::StringRef{"AND", 3});
1025+
and_node->add_child(const_cast<sql_parser::AstNode*>(lctx.where_expr));
1026+
and_node->add_child(const_cast<sql_parser::AstNode*>(rctx.where_expr));
1027+
where_expr = and_node;
1028+
} else if (lctx.where_expr) {
1029+
where_expr = lctx.where_expr;
1030+
} else {
1031+
where_expr = rctx.where_expr;
1032+
}
1033+
1034+
const auto& shard_list = shards_.get_shards(left_table->table_name);
1035+
PlanNode* current = nullptr;
1036+
for (const auto& shard : shard_list) {
1037+
sql_parser::StringRef sql = qb_.build_select_join(
1038+
left_table, right_table, join_node->join.condition, where_expr);
1039+
PlanNode* rs = make_remote_scan(shard.backend_name.c_str(), sql, left_table);
1040+
if (!current) {
1041+
current = rs;
1042+
} else {
1043+
PlanNode* union_node = make_plan_node(arena_, PlanNodeType::SET_OP);
1044+
union_node->set_op.op = SET_OP_UNION;
1045+
union_node->set_op.all = true;
1046+
union_node->left = current;
1047+
union_node->right = rs;
1048+
current = union_node;
1049+
}
1050+
}
1051+
return current ? current : join_node;
1052+
}
1053+
9471054
PlanNode* distribute_join(PlanNode* join_node) {
948-
// Get tables from each side
9491055
const TableInfo* left_table = find_table(join_node->left);
9501056
const TableInfo* right_table = find_table(join_node->right);
9511057

1058+
if (left_table && right_table &&
1059+
shards_.is_sharded(left_table->table_name) &&
1060+
shards_.is_sharded(right_table->table_name) &&
1061+
shards_.same_routing(left_table->table_name, right_table->table_name) &&
1062+
join_on_shard_keys(join_node->join.condition,
1063+
shards_.get_shard_key(left_table->table_name),
1064+
shards_.get_shard_key(right_table->table_name))) {
1065+
return distribute_colocated_join(join_node, left_table, right_table);
1066+
}
1067+
9521068
PlanNode* left_dist = nullptr;
9531069
PlanNode* right_dist = nullptr;
9541070

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

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

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

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

1283+
if (!table || !shards_.has_table(table->table_name)) return plan;
1284+
11791285
// Check for cross-shard subqueries in WHERE and rewrite
11801286
const sql_parser::AstNode* where_expr = dp.where_expr;
11811287
if (where_expr && has_subquery(where_expr) && remote_executor_) {
@@ -1262,7 +1368,68 @@ class DistributedPlanner {
12621368
return false;
12631369
}
12641370

1265-
// Scatter DML SQL to all shards, combining results via UNION ALL
1371+
void collect_ast_table_names(const sql_parser::AstNode* n,
1372+
std::vector<sql_parser::StringRef>& out) const {
1373+
if (!n) return;
1374+
if (n->type == sql_parser::NodeType::NODE_TABLE_REF && n->first_child) {
1375+
const sql_parser::AstNode* name = n->first_child;
1376+
if (name->type == sql_parser::NodeType::NODE_IDENTIFIER) {
1377+
out.push_back(name->value());
1378+
} else if (name->type == sql_parser::NodeType::NODE_QUALIFIED_NAME) {
1379+
const sql_parser::AstNode* schema = name->first_child;
1380+
const sql_parser::AstNode* table = schema ? schema->next_sibling : nullptr;
1381+
if (table) out.push_back(table->value());
1382+
else if (schema) out.push_back(schema->value());
1383+
}
1384+
}
1385+
for (const sql_parser::AstNode* c = n->first_child; c; c = c->next_sibling) {
1386+
collect_ast_table_names(c, out);
1387+
}
1388+
}
1389+
1390+
PlanNode* distribute_multi_table_dml(const sql_parser::AstNode* ast,
1391+
const TableInfo* primary,
1392+
bool is_update) {
1393+
std::vector<sql_parser::StringRef> names;
1394+
collect_ast_table_names(ast, names);
1395+
const char* backend = nullptr;
1396+
bool saw_mapped = false;
1397+
for (sql_parser::StringRef name : names) {
1398+
if (!shards_.has_table(name)) continue;
1399+
saw_mapped = true;
1400+
if (shards_.is_sharded(name)) {
1401+
return fail_dml(is_update
1402+
? "multi-table UPDATE is not supported on sharded tables"
1403+
: "multi-table DELETE is not supported on sharded tables");
1404+
}
1405+
const char* b = shards_.get_backend(name);
1406+
if (backend && b && std::strcmp(backend, b) != 0) {
1407+
return fail_dml(is_update
1408+
? "multi-table UPDATE spans multiple backends"
1409+
: "multi-table DELETE spans multiple backends");
1410+
}
1411+
if (b) backend = b;
1412+
}
1413+
if (!backend && primary && shards_.has_table(primary->table_name)) {
1414+
if (shards_.is_sharded(primary->table_name)) {
1415+
return fail_dml(is_update
1416+
? "multi-table UPDATE is not supported on sharded tables"
1417+
: "multi-table DELETE is not supported on sharded tables");
1418+
}
1419+
backend = shards_.get_backend(primary->table_name);
1420+
saw_mapped = true;
1421+
}
1422+
if (!backend || !saw_mapped) {
1423+
return fail_dml(is_update
1424+
? "multi-table UPDATE is not supported on sharded tables"
1425+
: "multi-table DELETE is not supported on sharded tables");
1426+
}
1427+
sql_parser::StringRef sql = is_update
1428+
? qb_.build_update_from_ast(ast)
1429+
: qb_.build_delete_from_ast(ast);
1430+
return make_remote_scan(backend, sql, primary);
1431+
}
1432+
12661433
PlanNode* scatter_dml_to_shards(const TableInfo* table,
12671434
const std::vector<ShardInfo>& shard_list,
12681435
std::function<sql_parser::StringRef()> build_sql) {

include/sql_engine/distributed_txn.h

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -178,6 +178,18 @@ class DistributedTransactionManager : public TransactionManager {
178178
return executor_.execute_dml(backend_name, sql);
179179
}
180180

181+
ResultSet route_query(const char* backend_name,
182+
sql_parser::StringRef sql) override {
183+
if (!active_) return executor_.execute(backend_name, sql);
184+
auto it = sessions_.find(backend_name);
185+
if (it != sessions_.end() && it->second) {
186+
return it->second->execute(sql);
187+
}
188+
return executor_.execute(backend_name, sql);
189+
}
190+
191+
bool route_query_supported() const override { return true; }
192+
181193
bool commit() override {
182194
if (!active_) return false;
183195
if (participants_.empty()) {

0 commit comments

Comments
 (0)