@@ -999,12 +999,72 @@ class DistributedPlanner {
999999 return result;
10001000 }
10011001
1002- // 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+
10031054 PlanNode* distribute_join (PlanNode* join_node) {
1004- // Get tables from each side
10051055 const TableInfo* left_table = find_table (join_node->left );
10061056 const TableInfo* right_table = find_table (join_node->right );
10071057
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+
10081068 PlanNode* left_dist = nullptr ;
10091069 PlanNode* right_dist = nullptr ;
10101070
@@ -1172,18 +1232,13 @@ class DistributedPlanner {
11721232 PlanNode* distribute_update (PlanNode* plan) {
11731233 const auto & up = plan->update_plan ;
11741234 const TableInfo* table = up.table ;
1175- if (!table || !shards_.has_table (table->table_name )) return plan;
11761235
1177- // Multi-table UPDATE: emit full SQL from AST, route to primary table's backend
11781236 if (up.original_ast ) {
1179- sql_parser::StringRef sql = qb_.build_update_from_ast (up.original_ast );
1180- if (!shards_.is_sharded (table->table_name )) {
1181- return make_remote_scan (shards_.get_backend (table->table_name ), sql, table);
1182- }
1183- const auto & shard_list = shards_.get_shards (table->table_name );
1184- return scatter_dml_to_shards (table, shard_list, [&]() { return sql; });
1237+ return distribute_multi_table_dml (up.original_ast , table, true );
11851238 }
11861239
1240+ if (!table || !shards_.has_table (table->table_name )) return plan;
1241+
11871242 // Check for cross-shard subqueries in WHERE and rewrite
11881243 const sql_parser::AstNode* where_expr = up.where_expr ;
11891244 if (where_expr && has_subquery (where_expr) && remote_executor_) {
@@ -1220,18 +1275,13 @@ class DistributedPlanner {
12201275 PlanNode* distribute_delete (PlanNode* plan) {
12211276 const auto & dp = plan->delete_plan ;
12221277 const TableInfo* table = dp.table ;
1223- if (!table || !shards_.has_table (table->table_name )) return plan;
12241278
1225- // Multi-table DELETE: emit full SQL from AST, route to primary table's backend
12261279 if (dp.original_ast ) {
1227- sql_parser::StringRef sql = qb_.build_delete_from_ast (dp.original_ast );
1228- if (!shards_.is_sharded (table->table_name )) {
1229- return make_remote_scan (shards_.get_backend (table->table_name ), sql, table);
1230- }
1231- const auto & shard_list = shards_.get_shards (table->table_name );
1232- return scatter_dml_to_shards (table, shard_list, [&]() { return sql; });
1280+ return distribute_multi_table_dml (dp.original_ast , table, false );
12331281 }
12341282
1283+ if (!table || !shards_.has_table (table->table_name )) return plan;
1284+
12351285 // Check for cross-shard subqueries in WHERE and rewrite
12361286 const sql_parser::AstNode* where_expr = dp.where_expr ;
12371287 if (where_expr && has_subquery (where_expr) && remote_executor_) {
@@ -1318,7 +1368,68 @@ class DistributedPlanner {
13181368 return false ;
13191369 }
13201370
1321- // 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+
13221433 PlanNode* scatter_dml_to_shards (const TableInfo* table,
13231434 const std::vector<ShardInfo>& shard_list,
13241435 std::function<sql_parser::StringRef()> build_sql) {
0 commit comments