@@ -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) {
0 commit comments