Skip to content

Commit e742dfd

Browse files
committed
fix: distribute derived tables and stop silent sort/SQL bugs
Rewrite DERIVED_SCAN inner plans so FROM (SELECT ...) hits remote shards. Store remote_sql_len as uint32_t so statements longer than 64KB are not truncated. Resolve ORDER BY position/alias in the plan builder, and only MERGE_SORT when every key is a table column — expressions gather and sort locally instead of comparing column 0.
1 parent 5da29e2 commit e742dfd

6 files changed

Lines changed: 236 additions & 11 deletions

File tree

include/sql_engine/distributed_planner.h

Lines changed: 63 additions & 7 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(

include/sql_engine/plan_builder.h

Lines changed: 54 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
#include "sql_parser/common.h"
2626
#include "sql_parser/arena.h"
2727
#include <cstring>
28+
#include <cstdlib>
2829
#include <vector>
2930

3031
namespace sql_engine {
@@ -116,6 +117,55 @@ class PlanBuilder {
116117
return false;
117118
}
118119

120+
static const sql_parser::AstNode* select_item_expr(const sql_parser::AstNode* item) {
121+
return item ? item->first_child : nullptr;
122+
}
123+
124+
static sql_parser::StringRef select_item_alias(const sql_parser::AstNode* item) {
125+
if (!item) return {};
126+
for (const sql_parser::AstNode* c = item->first_child; c; c = c->next_sibling) {
127+
if (c->type == sql_parser::NodeType::NODE_ALIAS) return c->value();
128+
}
129+
return {};
130+
}
131+
132+
static const sql_parser::AstNode* resolve_order_key(
133+
const sql_parser::AstNode* key, const sql_parser::AstNode* select_items) {
134+
if (!key || !select_items) return key;
135+
uint16_t n = count_children(select_items);
136+
if (n == 0) return key;
137+
138+
if (key->type == sql_parser::NodeType::NODE_LITERAL_INT) {
139+
sql_parser::StringRef sv = key->value();
140+
if (!sv.ptr || sv.len == 0) return key;
141+
int64_t pos = std::strtoll(sv.ptr, nullptr, 10);
142+
if (pos < 1 || pos > static_cast<int64_t>(n)) return key;
143+
uint16_t idx = 0;
144+
for (const sql_parser::AstNode* item = select_items->first_child; item;
145+
item = item->next_sibling, ++idx) {
146+
if (idx + 1 == static_cast<uint16_t>(pos)) {
147+
const sql_parser::AstNode* expr = select_item_expr(item);
148+
return expr ? expr : key;
149+
}
150+
}
151+
return key;
152+
}
153+
154+
if (key->type == sql_parser::NodeType::NODE_COLUMN_REF ||
155+
key->type == sql_parser::NodeType::NODE_IDENTIFIER) {
156+
sql_parser::StringRef name = key->value();
157+
for (const sql_parser::AstNode* item = select_items->first_child; item;
158+
item = item->next_sibling) {
159+
sql_parser::StringRef alias = select_item_alias(item);
160+
if (alias.ptr && alias.equals_ci(name.ptr, name.len)) {
161+
const sql_parser::AstNode* expr = select_item_expr(item);
162+
return expr ? expr : key;
163+
}
164+
}
165+
}
166+
return key;
167+
}
168+
119169
// Check if an expression (or any descendant) contains an aggregate function call.
120170
// Does NOT recurse into subqueries -- aggregates inside subqueries belong
121171
// to the subquery's own aggregation, not the outer query.
@@ -279,10 +329,12 @@ class PlanBuilder {
279329
arena_.allocate(sizeof(sql_parser::AstNode*) * cnt));
280330
auto* dirs = static_cast<uint8_t*>(arena_.allocate(cnt));
281331

332+
const sql_parser::AstNode* select_items =
333+
find_child(select_ast, sql_parser::NodeType::NODE_SELECT_ITEM_LIST);
334+
282335
uint16_t idx = 0;
283336
for (const sql_parser::AstNode* item = order_by->first_child; item; item = item->next_sibling) {
284-
// First child is the key expression
285-
keys[idx] = item->first_child;
337+
keys[idx] = resolve_order_key(item->first_child, select_items);
286338
// Check for DESC direction (second child with "DESC" value)
287339
dirs[idx] = 0; // ASC by default
288340
const sql_parser::AstNode* dir_node = find_child(item, sql_parser::NodeType::NODE_IDENTIFIER);

include/sql_engine/plan_executor.h

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,7 @@
5656
#include <string>
5757
#include <vector>
5858
#include <memory>
59+
#include <cstdlib>
5960

6061
namespace sql_engine {
6162

@@ -1170,12 +1171,21 @@ class PlanExecutor {
11701171

11711172
uint16_t resolve_column_index(const sql_parser::AstNode* key, const TableInfo* table) {
11721173
if (!key || !table) return 0;
1174+
if (key->type == sql_parser::NodeType::NODE_LITERAL_INT) {
1175+
sql_parser::StringRef sv = key->value();
1176+
if (!sv.ptr || sv.len == 0) return 0;
1177+
int64_t n = std::strtoll(sv.ptr, nullptr, 10);
1178+
if (n < 1) return 0;
1179+
if (n > static_cast<int64_t>(table->column_count)) {
1180+
return static_cast<uint16_t>(table->column_count - 1);
1181+
}
1182+
return static_cast<uint16_t>(n - 1);
1183+
}
11731184
sql_parser::StringRef col_name;
11741185
if (key->type == sql_parser::NodeType::NODE_COLUMN_REF ||
11751186
key->type == sql_parser::NodeType::NODE_IDENTIFIER) {
11761187
col_name = key->value();
11771188
} else if (key->type == sql_parser::NodeType::NODE_QUALIFIED_NAME) {
1178-
// table.column -- get the column part
11791189
const sql_parser::AstNode* c = key->first_child;
11801190
if (c && c->next_sibling) col_name = c->next_sibling->value();
11811191
else if (c) col_name = c->value();

include/sql_engine/plan_node.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,7 @@ struct PlanNode {
101101
struct {
102102
const char* backend_name;
103103
const char* remote_sql;
104-
uint16_t remote_sql_len;
104+
uint32_t remote_sql_len;
105105
const TableInfo* table; // expected result schema (for SELECT *)
106106
// Optional projection expressions used to derive result column
107107
// names when the remote SQL is not a passthrough SELECT *. When

tests/test_distributed_planner.cpp

Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -390,6 +390,10 @@ class DistributedPlannerTest : public ::testing::Test {
390390
find_nodes(node->merge_sort.children[i], type, out);
391391
return;
392392
}
393+
if (node->type == PlanNodeType::DERIVED_SCAN) {
394+
find_nodes(node->derived_scan.inner_plan, type, out);
395+
return;
396+
}
393397
find_nodes(node->left, type, out);
394398
find_nodes(node->right, type, out);
395399
}
@@ -964,3 +968,96 @@ TEST_F(DistributedPlannerTest, WindowGatherCorrectness) {
964968
EXPECT_EQ(dist_rs.row_count(), 15u);
965969
EXPECT_TRUE(compare_results_unordered(local_rs, dist_rs));
966970
}
971+
972+
TEST_F(DistributedPlannerTest, DerivedScanIsDistributed) {
973+
Parser<Dialect::MySQL> parser;
974+
const char* sql = "SELECT name FROM (SELECT name, age FROM users WHERE age > 20) AS t";
975+
auto pr = parser.parse(sql, std::strlen(sql));
976+
ASSERT_EQ(pr.status, ParseResult::OK);
977+
978+
PlanBuilder<Dialect::MySQL> builder(catalog, parser.arena());
979+
PlanNode* plan = builder.build(pr.ast);
980+
ASSERT_NE(plan, nullptr);
981+
982+
DistributedPlanner<Dialect::MySQL> dp(shard_map, catalog, parser.arena());
983+
PlanNode* dist = dp.distribute(plan);
984+
ASSERT_NE(dist, nullptr);
985+
986+
std::vector<PlanNode*> scans, remotes;
987+
find_nodes(dist, PlanNodeType::SCAN, scans);
988+
find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes);
989+
EXPECT_TRUE(scans.empty()) << "inner SCAN must be rewritten";
990+
EXPECT_FALSE(remotes.empty());
991+
}
992+
993+
TEST_F(DistributedPlannerTest, DerivedScanCorrectness) {
994+
const char* sql = "SELECT name FROM (SELECT name, age FROM users WHERE age > 20) AS t";
995+
auto local_rs = execute_local(sql);
996+
auto dist_rs = execute_distributed(sql);
997+
EXPECT_GT(local_rs.row_count(), 0u);
998+
EXPECT_EQ(local_rs.row_count(), dist_rs.row_count());
999+
EXPECT_TRUE(compare_results_unordered(local_rs, dist_rs));
1000+
}
1001+
1002+
TEST_F(DistributedPlannerTest, OrderByExpressionDoesNotMergeOnColumnZero) {
1003+
Parser<Dialect::MySQL> parser;
1004+
const char* sql = "SELECT name, age FROM users ORDER BY age + 1";
1005+
auto pr = parser.parse(sql, std::strlen(sql));
1006+
ASSERT_EQ(pr.status, ParseResult::OK);
1007+
1008+
PlanBuilder<Dialect::MySQL> builder(catalog, parser.arena());
1009+
PlanNode* plan = builder.build(pr.ast);
1010+
ASSERT_NE(plan, nullptr);
1011+
1012+
DistributedPlanner<Dialect::MySQL> dp(shard_map, catalog, parser.arena());
1013+
PlanNode* dist = dp.distribute(plan);
1014+
ASSERT_NE(dist, nullptr);
1015+
1016+
std::vector<PlanNode*> merges;
1017+
find_nodes(dist, PlanNodeType::MERGE_SORT, merges);
1018+
EXPECT_TRUE(merges.empty()) << "expression ORDER BY must not MERGE_SORT on col 0";
1019+
1020+
auto local_rs = execute_local(sql);
1021+
auto dist_rs = execute_distributed(sql);
1022+
EXPECT_EQ(local_rs.row_count(), dist_rs.row_count());
1023+
EXPECT_TRUE(compare_results_ordered(local_rs, dist_rs));
1024+
}
1025+
1026+
TEST_F(DistributedPlannerTest, OrderByAliasAndPositionCorrectness) {
1027+
const char* sql = "SELECT name AS n, age AS a FROM users ORDER BY a DESC";
1028+
auto local_rs = execute_local(sql);
1029+
auto dist_rs = execute_distributed(sql);
1030+
EXPECT_EQ(local_rs.row_count(), 15u);
1031+
EXPECT_TRUE(compare_results_ordered(local_rs, dist_rs));
1032+
1033+
const char* sql2 = "SELECT name, age FROM users ORDER BY 2 DESC";
1034+
auto local2 = execute_local(sql2);
1035+
auto dist2 = execute_distributed(sql2);
1036+
EXPECT_TRUE(compare_results_ordered(local2, dist2));
1037+
}
1038+
1039+
TEST_F(DistributedPlannerTest, RemoteSqlLenIsNotTruncated) {
1040+
std::string name(70000, 'x');
1041+
std::string sql = "SELECT * FROM users WHERE name = '" + name + "'";
1042+
Parser<Dialect::MySQL> parser;
1043+
auto pr = parser.parse(sql.c_str(), sql.size());
1044+
ASSERT_EQ(pr.status, ParseResult::OK);
1045+
1046+
PlanBuilder<Dialect::MySQL> builder(catalog, parser.arena());
1047+
PlanNode* plan = builder.build(pr.ast);
1048+
ASSERT_NE(plan, nullptr);
1049+
1050+
DistributedPlanner<Dialect::MySQL> dp(shard_map, catalog, parser.arena());
1051+
PlanNode* dist = dp.distribute(plan);
1052+
ASSERT_NE(dist, nullptr);
1053+
1054+
std::vector<PlanNode*> remotes;
1055+
find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes);
1056+
ASSERT_FALSE(remotes.empty());
1057+
for (auto* rs : remotes) {
1058+
EXPECT_GT(rs->remote_scan.remote_sql_len, 65535u);
1059+
std::string remote(rs->remote_scan.remote_sql, rs->remote_scan.remote_sql_len);
1060+
EXPECT_NE(remote.find(name), std::string::npos);
1061+
EXPECT_EQ(rs->remote_scan.remote_sql_len, remote.size());
1062+
}
1063+
}

tests/test_plan_executor.cpp

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -147,6 +147,16 @@ TEST_F(PlanExecutorTest, SelectDistinctDept) {
147147
EXPECT_EQ(depts.size(), 2u);
148148
}
149149

150+
TEST_F(PlanExecutorTest, OrderByAliasAndPosition) {
151+
auto by_alias = run_query("SELECT name AS n, age AS a FROM users ORDER BY a DESC");
152+
ASSERT_EQ(by_alias.row_count(), 5u);
153+
EXPECT_EQ(by_alias.rows[0].get(1).int_val, 35);
154+
155+
auto by_pos = run_query("SELECT name, age FROM users ORDER BY 2 DESC");
156+
ASSERT_EQ(by_pos.row_count(), 5u);
157+
EXPECT_EQ(by_pos.rows[0].get(1).int_val, 35);
158+
}
159+
150160
TEST_F(PlanExecutorTest, CountDistinctDept) {
151161
parser.reset();
152162
const char* sql = "SELECT COUNT(DISTINCT dept) FROM users";

0 commit comments

Comments
 (0)