diff --git a/Makefile b/Makefile index 984e413..c107ebb 100644 --- a/Makefile +++ b/Makefile @@ -88,7 +88,8 @@ TEST_SRCS = $(TEST_DIR)/test_main.cpp \ $(TEST_DIR)/test_result_set.cpp \ $(TEST_DIR)/test_ssl_config.cpp \ $(TEST_DIR)/test_star_modifiers.cpp \ - $(TEST_DIR)/test_shard_map.cpp + $(TEST_DIR)/test_shard_map.cpp \ + $(TEST_DIR)/test_shard_battery.cpp TEST_OBJS = $(TEST_SRCS:.cpp=.o) TEST_TARGET = $(PROJECT_ROOT)/run_tests diff --git a/include/sql_engine/distributed_planner.h b/include/sql_engine/distributed_planner.h index d224065..7e415fc 100644 --- a/include/sql_engine/distributed_planner.h +++ b/include/sql_engine/distributed_planner.h @@ -25,6 +25,7 @@ #include #include #include +#include #include namespace sql_engine { @@ -78,12 +79,19 @@ class DistributedPlanner { RemoteExecutor* remote_executor_; FunctionRegistry* functions_; const char* error_; + std::string error_storage_; PlanNode* fail_dml(const char* message) { error_ = message; return nullptr; } + PlanNode* fail_dml_owned(std::string message) { + error_storage_ = std::move(message); + error_ = error_storage_.c_str(); + return nullptr; + } + // Push aggregate expressions from PROJECT into AGGREGATE node // (same logic as PlanExecutor::preprocess_aggregates) void push_agg_exprs_from_project(PlanNode* project_node, PlanNode* agg_node) { @@ -159,8 +167,11 @@ class DistributedPlanner { agg_child = agg_child->left; } if (agg_child && agg_child->type == PlanNodeType::AGGREGATE) { - push_agg_exprs_from_project(node, agg_child); - PlanNode* dist_agg = distribute_aggregate(agg_child); + PlanNode* agg_copy = make_plan_node(arena_, PlanNodeType::AGGREGATE); + agg_copy->aggregate = agg_child->aggregate; + agg_copy->left = agg_child->left; + push_agg_exprs_from_project(node, agg_copy); + PlanNode* dist_agg = distribute_aggregate(agg_copy); if (dist_agg && (dist_agg->type == PlanNodeType::MERGE_AGGREGATE || dist_agg->type == PlanNodeType::AGGREGATE)) { PlanNode* top = dist_agg; @@ -1193,7 +1204,8 @@ class DistributedPlanner { 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); + left_table, right_table, join_node->join.condition, where_expr, + join_node->join.join_type); PlanNode* rs = make_remote_scan(shard.backend_name.c_str(), sql, left_table); if (!current) { current = rs; @@ -1310,6 +1322,7 @@ class DistributedPlanner { const TableInfo* right_table) { if (!join_node || !remote_executor_ || !join_node->join.condition) return nullptr; + if (join_node->join.join_type != JOIN_INNER) return nullptr; if (!left_table || !right_table) return nullptr; bool ls = shards_.is_sharded(left_table->table_name); @@ -1550,6 +1563,7 @@ class DistributedPlanner { const sql_parser::AstNode* where_expr = up.where_expr; if (where_expr && has_subquery(where_expr) && remote_executor_) { where_expr = rewrite_where_subquery(where_expr, table); + if (error_) return nullptr; } if (!shards_.is_sharded(table->table_name)) { @@ -1594,6 +1608,7 @@ class DistributedPlanner { const sql_parser::AstNode* where_expr = dp.where_expr; if (where_expr && has_subquery(where_expr) && remote_executor_) { where_expr = rewrite_where_subquery(where_expr, table); + if (error_) return nullptr; } if (!shards_.is_sharded(table->table_name)) { @@ -1659,7 +1674,7 @@ class DistributedPlanner { sql_parser::StringRef shard_key) const { if (!set_columns || !shard_key.ptr) return false; for (uint16_t i = 0; i < set_count; ++i) { - if (is_column_ref(set_columns[i], shard_key)) return true; + if (is_shard_key_ref(set_columns[i], shard_key)) return true; } return false; } @@ -1844,7 +1859,7 @@ class DistributedPlanner { auto resolve = make_resolver(catalog_, table, src.values); for (uint16_t i = 0; i < set_count; ++i) { if (!set_cols[i]) continue; - const ColumnInfo* col = catalog_.get_column(table, set_cols[i]->value()); + const ColumnInfo* col = catalog_.get_column(table, set_col_name(set_cols[i])); if (!col) continue; Value nv = value_null(); if (functions_) { @@ -1935,8 +1950,16 @@ class DistributedPlanner { } sql_parser::StringRef sql = qb_.build_select( table, where_expr, nullptr, 0, nullptr, 0, - nullptr, nullptr, 0, -1, false); + nullptr, nullptr, 0, -1, false, true); ResultSet rs = remote_executor_->execute(shard.backend_name.c_str(), sql); + if (!rs.ok) { + std::string msg = rs.error_message.empty() + ? "shard-key UPDATE SELECT failed" : rs.error_message; + msg += " ["; + msg.append(sql.ptr, sql.len); + msg += "]"; + return fail_dml_owned(std::move(msg)); + } for (const auto& row : rs.rows) { Move m; m.src = src; @@ -1950,9 +1973,7 @@ class DistributedPlanner { } if (moves.empty()) { - sql_parser::StringRef sql = qb_.build_update( - table, up.set_columns, up.set_exprs, up.set_count, where_expr); - return make_remote_scan(pruned[0].backend_name.c_str(), sql, table); + return make_noop_update(table, pruned[0].backend_name.c_str()); } bool any_move = false; @@ -1999,6 +2020,31 @@ class DistributedPlanner { return current ? current : plan; } + static sql_parser::StringRef set_col_name(const sql_parser::AstNode* node) { + if (!node) return sql_parser::StringRef{nullptr, 0}; + if (node->type == sql_parser::NodeType::NODE_QUALIFIED_NAME) { + const sql_parser::AstNode* c = node->first_child; + if (c && c->next_sibling) return c->next_sibling->value(); + } + return node->value(); + } + + PlanNode* make_noop_update(const TableInfo* table, const char* backend) { + sql_parser::StringBuilder sb(arena_, 64); + sb.append("UPDATE "); + if (table) sb.append(table->table_name.ptr, table->table_name.len); + sb.append(" SET "); + if (table && table->column_count > 0) { + sb.append(table->columns[0].name.ptr, table->columns[0].name.len); + sb.append(" = "); + sb.append(table->columns[0].name.ptr, table->columns[0].name.len); + } else { + sb.append("id = id"); + } + sb.append(" WHERE 1 = 0"); + return make_remote_scan(backend, sb.finish(), table); + } + bool is_column_ref(const sql_parser::AstNode* node, sql_parser::StringRef col_name) const { if (!node) return false; if (node->type == sql_parser::NodeType::NODE_COLUMN_REF || @@ -2172,10 +2218,15 @@ class DistributedPlanner { // Execute: if it's a RemoteScan, execute via remote executor // Otherwise, need to execute locally ResultSet rs = execute_distributed_plan(dist_plan); + if (!rs.ok) { + error_storage_ = rs.error_message.empty() ? "subquery failed" : rs.error_message; + error_ = error_storage_.c_str(); + return result; + } for (const auto& row : rs.rows) { if (row.column_count > 0) { - result.push_back(row.get(0)); + result.push_back(copy_value_arena(row.get(0))); } } return result; @@ -2183,7 +2234,8 @@ class DistributedPlanner { // Execute a distributed plan tree (recursively handles SET_OP / REMOTE_SCAN). ResultSet execute_distributed_plan(PlanNode* node) { - if (!node || !remote_executor_) return {}; + if (!remote_executor_) return ResultSet::fail("no remote executor"); + if (!node) return ResultSet::fail("empty distributed plan"); if (node->type == PlanNodeType::REMOTE_SCAN) { sql_parser::StringRef sql{node->remote_scan.remote_sql, @@ -2192,9 +2244,10 @@ class DistributedPlanner { } if (node->type == PlanNodeType::SET_OP) { - // UNION ALL: concatenate results ResultSet left = execute_distributed_plan(node->left); + if (!left.ok) return left; ResultSet right = execute_distributed_plan(node->right); + if (!right.ok) return right; for (auto& row : right.rows) { left.rows.push_back(row); } @@ -2249,8 +2302,9 @@ class DistributedPlanner { lit = sql_parser::make_node(arena_, sql_parser::NodeType::NODE_LITERAL_INT, sql_parser::StringRef{s, static_cast(n)}); } else if (v.tag == Value::TAG_STRING && v.str_val.ptr) { + Value owned = copy_value_arena(v); lit = sql_parser::make_node(arena_, sql_parser::NodeType::NODE_LITERAL_STRING, - v.str_val); + owned.str_val); } else if (v.tag == Value::TAG_DOUBLE) { char buf[64]; int n = snprintf(buf, sizeof(buf), "%g", v.double_val); @@ -2269,6 +2323,19 @@ class DistributedPlanner { return new_in; } + sql_parser::AstNode* make_false_pred() { + sql_parser::AstNode* eq = sql_parser::make_node( + arena_, sql_parser::NodeType::NODE_BINARY_OP, + sql_parser::StringRef{"=", 1}); + eq->add_child(sql_parser::make_node( + arena_, sql_parser::NodeType::NODE_LITERAL_INT, + sql_parser::StringRef{"0", 1})); + eq->add_child(sql_parser::make_node( + arena_, sql_parser::NodeType::NODE_LITERAL_INT, + sql_parser::StringRef{"1", 1})); + return eq; + } + // Rewrite a WHERE expression by replacing the first IN (subquery) with IN (literals). // Returns the rewritten expression, or the original if no rewrite needed. const sql_parser::AstNode* rewrite_where_subquery( @@ -2286,7 +2353,7 @@ class DistributedPlanner { if (!values.empty()) { return build_in_list_from_values(where_expr, values); } - return where_expr; + return make_false_pred(); } } } @@ -2386,11 +2453,15 @@ class DistributedPlanner { PlanNode* dist_select = distribute_node(select_plan); ResultSet rs = execute_distributed_plan(dist_select); + if (!rs.ok) { + return fail_dml_owned(rs.error_message.empty() + ? "INSERT ... SELECT failed" : rs.error_message); + } if (rs.rows.empty()) { - // No rows to insert -- return a no-op - // Just return the original plan (which will do nothing since select_source is null) - return plan; + const auto& sl = shards_.get_shards(table->table_name); + if (sl.empty()) return fail_dml("table not in shard map"); + return make_noop_update(table, sl[0].backend_name.c_str()); } // Determine target shards for each row diff --git a/include/sql_engine/plan_executor.h b/include/sql_engine/plan_executor.h index 8bff423..8dcfe26 100644 --- a/include/sql_engine/plan_executor.h +++ b/include/sql_engine/plan_executor.h @@ -309,9 +309,8 @@ class PlanExecutor { rs.column_count = rs.rows[0].column_count; } - // Build column names from plan build_column_names(plan, rs); - + rs.ok = true; return rs; } diff --git a/include/sql_engine/remote_query_builder.h b/include/sql_engine/remote_query_builder.h index ae99ad9..c1ef561 100644 --- a/include/sql_engine/remote_query_builder.h +++ b/include/sql_engine/remote_query_builder.h @@ -30,7 +30,8 @@ class RemoteQueryBuilder { uint8_t* order_dirs, uint16_t order_count, int64_t limit, // -1 = no limit - bool distinct) + bool distinct, + bool for_update = false) { sql_parser::StringBuilder sb(arena_, 512); @@ -91,6 +92,8 @@ class RemoteQueryBuilder { sb.append(buf, n); } + if (for_update) sb.append(" FOR UPDATE"); + return sb.finish(); } @@ -98,12 +101,16 @@ class RemoteQueryBuilder { const TableInfo* left, const TableInfo* right, const sql_parser::AstNode* on_expr, - const sql_parser::AstNode* where_expr) + const sql_parser::AstNode* where_expr, + uint8_t join_type = 0) { sql_parser::StringBuilder sb(arena_, 512); sb.append("SELECT * FROM "); if (left) sb.append(left->table_name.ptr, left->table_name.len); - sb.append(" JOIN "); + if (join_type == 1) sb.append(" LEFT JOIN "); + else if (join_type == 2) sb.append(" RIGHT JOIN "); + else if (join_type == 3) sb.append(" FULL JOIN "); + else sb.append(" JOIN "); if (right) sb.append(right->table_name.ptr, right->table_name.len); if (on_expr) { sb.append(" ON "); diff --git a/include/sql_engine/result_set.h b/include/sql_engine/result_set.h index 0479d42..69f2317 100644 --- a/include/sql_engine/result_set.h +++ b/include/sql_engine/result_set.h @@ -43,6 +43,9 @@ struct ResultSet { // extending their lifetime to match this ResultSet. std::vector> backing_lifetimes; + bool ok = true; + std::string error_message; + ResultSet() = default; ~ResultSet() { for (auto* arr : owned_value_arrays) ::operator delete(arr); @@ -55,8 +58,11 @@ struct ResultSet { column_count(o.column_count), owned_value_arrays(std::move(o.owned_value_arrays)), owned_strings(std::move(o.owned_strings)), - backing_lifetimes(std::move(o.backing_lifetimes)) { + backing_lifetimes(std::move(o.backing_lifetimes)), + ok(o.ok), + error_message(std::move(o.error_message)) { o.column_count = 0; + o.ok = true; } ResultSet& operator=(ResultSet&& o) noexcept { @@ -68,7 +74,10 @@ struct ResultSet { owned_value_arrays = std::move(o.owned_value_arrays); owned_strings = std::move(o.owned_strings); backing_lifetimes = std::move(o.backing_lifetimes); + ok = o.ok; + error_message = std::move(o.error_message); o.column_count = 0; + o.ok = true; } return *this; } @@ -80,6 +89,13 @@ struct ResultSet { size_t row_count() const { return rows.size(); } bool empty() const { return rows.empty(); } + static ResultSet fail(const char* msg) { + ResultSet rs; + rs.ok = false; + rs.error_message = msg ? msg : "query failed"; + return rs; + } + // Allocate a heap-owned row and append it to rows. Returns a reference // to the Row (which points into owned_value_arrays). Row& add_heap_row(uint16_t col_count) { diff --git a/include/sql_engine/session.h b/include/sql_engine/session.h index 26d2494..b4f6e72 100644 --- a/include/sql_engine/session.h +++ b/include/sql_engine/session.h @@ -132,7 +132,7 @@ class Session { exec_arena_.reset(); auto& entry = *cache_it->second; PlanNode* plan = maybe_distribute(entry.plan, exec_arena_); - if (!plan) return {}; + if (!plan) return ResultSet::fail(last_query_error_.c_str()); PlanExecutor executor(functions_, catalog_, exec_arena_); wire_executor(executor); return executor.execute(plan); @@ -169,7 +169,7 @@ class Session { if (shard_map_ && remote_executor_) { exec_arena_.reset(); PlanNode* dist = maybe_distribute(plan, exec_arena_); - if (!dist) return {}; + if (!dist) return ResultSet::fail(last_query_error_.c_str()); PlanExecutor executor(functions_, catalog_, exec_arena_); wire_executor(executor); rs = executor.execute(dist); @@ -256,8 +256,9 @@ class Session { // If sharding is configured, distribute DML to remote backends. if (shard_map_ && remote_executor_) { + routing_exec_.bind(remote_executor_, &txn_mgr_); DistributedPlanner dp(*shard_map_, catalog_, parser_.arena(), - remote_executor_, &functions_); + &routing_exec_, &functions_); PlanNode* dist_plan = dp.distribute_dml(plan); if (dp.last_error()) { @@ -274,10 +275,15 @@ class Session { dist_plan->remote_scan.backend_name, sql_ref); } } else if (dist_plan && dist_plan->type == PlanNodeType::SET_OP) { - // Scatter DML to multiple shards + if (plan_is_row_move(dist_plan) && !txn_mgr_.is_distributed()) { + result.success = false; + result.error_message = + "cross-shard row move requires a distributed transaction"; + } else { result.success = true; result.affected_rows = 0; for_each_remote_scan(dist_plan, [&](const PlanNode* rs) { + if (!result.success) return; sql_parser::StringRef s{rs->remote_scan.remote_sql, rs->remote_scan.remote_sql_len}; DmlResult shard_result; @@ -293,6 +299,7 @@ class Session { } result.affected_rows += shard_result.affected_rows; }); + } } else { // Not distributed (table not in shard map) -- local execution PlanExecutor executor(functions_, catalog_, parser_.arena()); @@ -374,13 +381,19 @@ class Session { CacheList plan_cache_order_; std::unordered_map plan_cache_; size_t plan_cache_max_size_ = 1024; + std::string last_query_error_; PlanNode* maybe_distribute(PlanNode* plan, sql_parser::Arena& arena) { + last_query_error_.clear(); if (!plan || !shard_map_ || !remote_executor_) return plan; + routing_exec_.bind(remote_executor_, &txn_mgr_); DistributedPlanner dplanner(*shard_map_, catalog_, arena, - remote_executor_, &functions_); + &routing_exec_, &functions_); PlanNode* dist = dplanner.distribute(plan); - if (dplanner.last_error()) return nullptr; + if (dplanner.last_error()) { + last_query_error_ = dplanner.last_error(); + return nullptr; + } return dist; } @@ -405,6 +418,18 @@ class Session { plan_cache_[plan_cache_order_.front().key] = plan_cache_order_.begin(); } + static bool plan_is_row_move(const PlanNode* node) { + bool has_del = false, has_ins = false; + for_each_remote_scan(node, [&](const PlanNode* rs) { + const char* s = rs->remote_scan.remote_sql; + uint32_t n = rs->remote_scan.remote_sql_len; + if (!s || n < 6) return; + if (n >= 6 && (s[0] == 'D' || s[0] == 'd')) has_del = true; + if (n >= 6 && (s[0] == 'I' || s[0] == 'i')) has_ins = true; + }); + return has_del && has_ins; + } + static void for_each_remote_scan(const PlanNode* node, const std::function& fn) { if (!node) return; diff --git a/include/sql_engine/shard_map.h b/include/sql_engine/shard_map.h index 894a5cd..e8ca528 100644 --- a/include/sql_engine/shard_map.h +++ b/include/sql_engine/shard_map.h @@ -9,6 +9,7 @@ #include #include #include +#include namespace sql_engine { @@ -43,7 +44,7 @@ struct ShardRange { // Used by LIST strategy. Each entry maps a single key value to a shard // index. Key may be int or string, but a single TableShardConfig must // stay one or the other (mixed lists are not supported). Lookups that -// miss every entry fall back to shard 0. +// miss every entry is unroutable (try_* returns false). struct ShardListEntry { bool is_int = true; int64_t int_val = 0; @@ -132,9 +133,17 @@ class ShardMap { case RoutingStrategy::HASH: out = fnv1a_int64(value) % n; return true; - case RoutingStrategy::RANGE: - out = shard_index_for_int(table_name, value); + case RoutingStrategy::RANGE: { + if (cfg->ranges.empty()) return false; + for (const auto& r : cfg->ranges) { + if (value <= r.upper_inclusive) { + out = clamp_index(r.shard_index, n); + return true; + } + } + out = clamp_index(cfg->ranges.back().shard_index, n); return true; + } case RoutingStrategy::LIST: for (const auto& e : cfg->list) { if (e.is_int && e.int_val == value) { @@ -154,9 +163,13 @@ class ShardMap { if (!cfg || cfg->shards.empty()) return false; size_t n = cfg->shards.size(); switch (cfg->strategy) { - case RoutingStrategy::HASH: + case RoutingStrategy::HASH: { + int64_t as_int = 0; + if (parse_full_int(val, val_len, as_int)) + return try_shard_index_for_int(table_name, as_int, out); out = fnv1a_bytes(reinterpret_cast(val), val_len) % n; return true; + } case RoutingStrategy::RANGE: return false; case RoutingStrategy::LIST: @@ -204,60 +217,19 @@ class ShardMap { return true; } - // Determine which shard index a value maps to. Dispatches on the - // configured RoutingStrategy. Returns 0 if the table is unknown or - // has no shards — prefer try_shard_index_for_* which fails closed. + // Prefer try_shard_index_for_*. These return SIZE_MAX if unroutable. size_t shard_index_for_int(sql_parser::StringRef table_name, int64_t value) const { - const TableShardConfig* cfg = lookup(table_name); - if (!cfg || cfg->shards.empty()) return 0; - size_t n = cfg->shards.size(); - switch (cfg->strategy) { - case RoutingStrategy::HASH: - return fnv1a_int64(value) % n; - case RoutingStrategy::RANGE: { - if (cfg->ranges.empty()) return 0; - for (const auto& r : cfg->ranges) { - if (value <= r.upper_inclusive) { - return clamp_index(r.shard_index, n); - } - } - // Above all upper bounds: fall through to last shard. - return clamp_index(cfg->ranges.back().shard_index, n); - } - case RoutingStrategy::LIST: - for (const auto& e : cfg->list) { - if (e.is_int && e.int_val == value) { - return clamp_index(e.shard_index, n); - } - } - return 0; - } - return 0; + size_t out = 0; + if (!try_shard_index_for_int(table_name, value, out)) return static_cast(-1); + return out; } size_t shard_index_for_string(sql_parser::StringRef table_name, const char* val, uint32_t val_len) const { - const TableShardConfig* cfg = lookup(table_name); - if (!cfg || cfg->shards.empty()) return 0; - size_t n = cfg->shards.size(); - switch (cfg->strategy) { - case RoutingStrategy::HASH: - return fnv1a_bytes(reinterpret_cast(val), val_len) % n; - case RoutingStrategy::RANGE: - // RANGE is integer-keyed only. Fall back to scatter-friendly - // shard 0 rather than producing a misleading single-shard - // route from a string key. - return 0; - case RoutingStrategy::LIST: - for (const auto& e : cfg->list) { - if (!e.is_int && e.str_val.size() == val_len && - std::memcmp(e.str_val.data(), val, val_len) == 0) { - return clamp_index(e.shard_index, n); - } - } - return 0; - } - return 0; + size_t out = 0; + if (!try_shard_index_for_string(table_name, val, val_len, out)) + return static_cast(-1); + return out; } RoutingStrategy routing_strategy(sql_parser::StringRef table_name) const { @@ -357,6 +329,18 @@ class ShardMap { return idx < n ? idx : (n == 0 ? 0 : n - 1); } + static bool parse_full_int(const char* val, uint32_t val_len, int64_t& out) { + if (!val || val_len == 0 || val_len > 20) return false; + char buf[24]; + std::memcpy(buf, val, val_len); + buf[val_len] = '\0'; + char* end = nullptr; + long long n = std::strtoll(buf, &end, 10); + if (!end || end != buf + val_len) return false; + out = static_cast(n); + return true; + } + static void split_keys(const std::string& spec, std::vector& out) { size_t start = 0; while (start <= spec.size()) { diff --git a/tests/test_distributed_dml.cpp b/tests/test_distributed_dml.cpp index 513afee..e7a090c 100644 --- a/tests/test_distributed_dml.cpp +++ b/tests/test_distributed_dml.cpp @@ -97,12 +97,16 @@ class DmlMockRemoteExecutor : public RemoteExecutor { ResultSet execute(const char* backend_name, StringRef sql) override { auto it = backends_.find(backend_name); - if (it == backends_.end()) return {}; + if (it == backends_.end()) return ResultSet::fail("unknown backend"); DmlBackendData* bd = it->second.get(); bd->executed_sqls.emplace_back(sql.ptr, sql.len); std::string sql_str(sql.ptr, sql.len); + const char* fu = " FOR UPDATE"; + if (sql_str.size() > 11 && + sql_str.compare(sql_str.size() - 11, 11, fu) == 0) + sql_str.resize(sql_str.size() - 11); // Detect DML vs SELECT if (is_dml(sql_str)) { @@ -124,7 +128,9 @@ class DmlMockRemoteExecutor : public RemoteExecutor { for (auto& [tname, src] : bd->mutable_sources) { executor.add_mutable_data_source(tname.c_str(), src); } - return executor.execute(plan); + ResultSet out = executor.execute(plan); + out.ok = true; + return out; } DmlResult execute_dml(const char* backend_name, StringRef sql) override { @@ -730,6 +736,36 @@ TEST_F(DistributedDmlTest, UpdateShardKeyMovesRow) { "Carol"); } +TEST_F(DistributedDmlTest, UpdateQualifiedShardKeyMovesRow) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + const char* src = backend_for_id(3); + const char* dst = backend_for_id(9); + ASSERT_STRNE(src, dst); + + auto result = execute_distributed_dml("UPDATE users SET users.id = 9 WHERE id = 3"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(row_count_on(src, "users"), 0u); + EXPECT_EQ(row_count_on(dst, "users"), 1u); +} + +TEST_F(DistributedDmlTest, UpdateShardKeyNoMatchingRowIsNoop) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + const char* home = backend_for_id(3); + auto result = execute_distributed_dml("UPDATE users SET id = 9 WHERE id = 99"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(row_count_on(home, "users"), 1u); + EXPECT_EQ(mock_executor.total_row_count("users"), 1u); +} + +TEST_F(DistributedDmlTest, InsertStringIntHashesLikeInt) { + auto ins = execute_distributed_dml( + "INSERT INTO users (id, name, age) VALUES ('3', 'Carol', 17)"); + EXPECT_TRUE(ins.success) << ins.error_message; + EXPECT_EQ(row_count_on(backend_for_id(3), "users"), 1u); + auto got = execute_distributed_select("SELECT name FROM users WHERE id = 3"); + ASSERT_EQ(got.row_count(), 1u); +} + TEST_F(DistributedDmlTest, UpdateShardKeySameShard) { int64_t a = 3; int64_t b = a; @@ -956,3 +992,176 @@ TEST_F(DistributedDmlTest, PlanCacheSeesUpdatedShardMap) { EXPECT_EQ(mock_executor.get_executed_sqls(s).size(), 0u) << s; } } + +TEST_F(DistributedDmlTest, UpdateQualifiedNonKeyDoesNotMove) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + const char* home = backend_for_id(3); + auto result = execute_distributed_dml("UPDATE users SET users.age = 40 WHERE id = 3"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(row_count_on(home, "users"), 1u); + EXPECT_EQ(mock_executor.total_row_count("users"), 1u); +} + +TEST_F(DistributedDmlTest, UpdateShardKeyToSameValue) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + auto result = execute_distributed_dml("UPDATE users SET id = 3 WHERE id = 3"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(row_count_on(backend_for_id(3), "users"), 1u); + EXPECT_EQ(mock_executor.total_row_count("users"), 1u); +} + +TEST_F(DistributedDmlTest, UpdateMovesThenPointSelectAndDelete) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + ASSERT_STRNE(backend_for_id(3), backend_for_id(9)); + EXPECT_TRUE(execute_distributed_dml("UPDATE users SET id = 9 WHERE id = 3").success); + EXPECT_EQ(execute_distributed_select("SELECT name FROM users WHERE id = 3").row_count(), 0u); + EXPECT_EQ(execute_distributed_select("SELECT name FROM users WHERE id = 9").row_count(), 1u); + EXPECT_TRUE(execute_distributed_dml("DELETE FROM users WHERE id = 9").success); + EXPECT_EQ(mock_executor.total_row_count("users"), 0u); +} + +TEST_F(DistributedDmlTest, UpdateMovesMultipleRowsToOneShard) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'A', 1)"); + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (4, 'B', 2)"); + EXPECT_TRUE(execute_distributed_dml("UPDATE users SET id = 9 WHERE id IN (3, 4)").success); + EXPECT_EQ(row_count_on(backend_for_id(9), "users"), 2u); + EXPECT_EQ(mock_executor.total_row_count("users"), 2u); +} + +TEST_F(DistributedDmlTest, InsertNegativeStringInt) { + auto ins = execute_distributed_dml( + "INSERT INTO users (id, name, age) VALUES ('-7', 'Neg', 1)"); + EXPECT_TRUE(ins.success) << ins.error_message; + EXPECT_EQ(row_count_on(backend_for_id(-7), "users"), 1u); + EXPECT_EQ(execute_distributed_select("SELECT name FROM users WHERE id = -7").row_count(), 1u); +} + +TEST_F(DistributedDmlTest, InsertLeadingZeroStringInt) { + auto ins = execute_distributed_dml( + "INSERT INTO users (id, name, age) VALUES ('03', 'Zed', 1)"); + EXPECT_TRUE(ins.success) << ins.error_message; + EXPECT_EQ(row_count_on(backend_for_id(3), "users"), 1u); + EXPECT_EQ(execute_distributed_select("SELECT name FROM users WHERE id = 3").row_count(), 1u); +} + +TEST_F(DistributedDmlTest, EmptyInSubqueryDeletesNothing) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + mock_executor.clear_sql_logs(); + auto result = execute_distributed_dml( + "DELETE FROM users WHERE id IN (SELECT user_id FROM orders)"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(mock_executor.total_row_count("users"), 1u); +} + +TEST_F(DistributedDmlTest, EmptyInSubqueryUpdateTouchesNothing) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + auto result = execute_distributed_dml( + "UPDATE users SET age = 99 WHERE id IN (SELECT user_id FROM orders)"); + EXPECT_TRUE(result.success) << result.error_message; + auto got = execute_distributed_select("SELECT age FROM users WHERE id = 3"); + ASSERT_EQ(got.row_count(), 1u); + EXPECT_EQ(got.rows[0].get(0).int_val, 17); +} + +TEST_F(DistributedDmlTest, InSubqueryStringNames) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (4, 'Dave', 18)"); + auto result = execute_distributed_dml( + "DELETE FROM users WHERE name IN (SELECT name FROM users WHERE id = 3)"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(mock_executor.total_row_count("users"), 1u); + EXPECT_EQ(execute_distributed_select("SELECT name FROM users WHERE id = 4").row_count(), 1u); +} + +TEST_F(DistributedDmlTest, InsertSelectEmptySourceIsNoop) { + auto result = execute_distributed_dml( + "INSERT INTO users (id, name, age) SELECT order_id, 'x', 1 FROM orders WHERE order_id = 999"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(mock_executor.total_row_count("users"), 0u); +} + +TEST_F(DistributedDmlTest, UnknownTableSelectIsEmpty) { + catalog.add_table("", "ghost", {{"id", SqlType::make_int(), false}}); + LocalTransactionManager txn(data_arena); + Session session(catalog, txn); + session.set_remote_executor(&mock_executor); + session.set_shard_map(&shard_map); + auto rs = session.execute_query("SELECT * FROM ghost"); + EXPECT_FALSE(rs.ok); + EXPECT_NE(rs.error_message.find("shard map"), std::string::npos); +} + +TEST_F(DistributedDmlTest, UpdateUnknownTableErrors) { + catalog.add_table("", "ghost", {{"id", SqlType::make_int(), false}}); + auto result = execute_distributed_dml("UPDATE ghost SET id = 1"); + EXPECT_FALSE(result.success); + EXPECT_NE(result.error_message.find("shard map"), std::string::npos); +} + +TEST_F(DistributedDmlTest, DeleteUnknownTableErrors) { + catalog.add_table("", "ghost", {{"id", SqlType::make_int(), false}}); + auto result = execute_distributed_dml("DELETE FROM ghost"); + EXPECT_FALSE(result.success); + EXPECT_NE(result.error_message.find("shard map"), std::string::npos); +} + +TEST_F(DistributedDmlTest, PlanCacheCountTwice) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (4, 'Dave', 18)"); + LocalTransactionManager txn(data_arena); + Session session(catalog, txn); + session.set_remote_executor(&mock_executor); + session.set_shard_map(&shard_map); + const char* sql = "SELECT COUNT(*) FROM users"; + auto a = session.execute_query(sql); + auto b = session.execute_query(sql); + ASSERT_EQ(a.row_count(), 1u); + ASSERT_EQ(b.row_count(), 1u); + EXPECT_EQ(a.rows[0].get(0).tag, b.rows[0].get(0).tag); + EXPECT_EQ(a.rows[0].get(0).to_int64(), 2); + EXPECT_EQ(b.rows[0].get(0).to_int64(), 2); + EXPECT_EQ(session.plan_cache_size(), 1u); +} + +TEST_F(DistributedDmlTest, PlanCacheSumAndGroupByTwice) { + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (3, 'Carol', 17)"); + execute_distributed_dml("INSERT INTO users (id, name, age) VALUES (4, 'Dave', 18)"); + LocalTransactionManager txn(data_arena); + Session session(catalog, txn); + session.set_remote_executor(&mock_executor); + session.set_shard_map(&shard_map); + const char* sql = "SELECT SUM(age) FROM users"; + auto a = session.execute_query(sql); + auto b = session.execute_query(sql); + ASSERT_EQ(a.row_count(), 1u); + ASSERT_EQ(b.row_count(), 1u); + EXPECT_EQ(a.rows[0].get(0).to_int64(), 35); + EXPECT_EQ(b.rows[0].get(0).to_int64(), 35); +} + +TEST_F(DistributedDmlTest, CompositeQualifiedShardKeyMove) { + catalog.add_table("", "kv", { + {"tenant_id", SqlType::make_int(), false}, + {"id", SqlType::make_int(), false}, + {"name", SqlType::make_varchar(255), true}, + }); + TableShardConfig cfg; + cfg.table_name = "kv"; + cfg.shard_key = "tenant_id+id"; + cfg.shards = {{"shard0"}, {"shard1"}, {"shard2"}}; + shard_map.add_table(cfg); + mock_executor.add_table_to_all("kv", { + {"tenant_id", SqlType::make_int(), false}, + {"id", SqlType::make_int(), false}, + {"name", SqlType::make_varchar(255), true}, + }); + EXPECT_TRUE(execute_distributed_dml( + "INSERT INTO kv (tenant_id, id, name) VALUES (1, 3, 'A')").success); + auto result = execute_distributed_dml( + "UPDATE kv SET kv.id = 9 WHERE tenant_id = 1 AND id = 3"); + EXPECT_TRUE(result.success) << result.error_message; + EXPECT_EQ(execute_distributed_select( + "SELECT name FROM kv WHERE tenant_id = 1 AND id = 9").row_count(), 1u); + EXPECT_EQ(execute_distributed_select( + "SELECT name FROM kv WHERE tenant_id = 1 AND id = 3").row_count(), 0u); +} diff --git a/tests/test_distributed_planner.cpp b/tests/test_distributed_planner.cpp index f5f1c88..da4d14a 100644 --- a/tests/test_distributed_planner.cpp +++ b/tests/test_distributed_planner.cpp @@ -1200,6 +1200,51 @@ TEST_F(DistributedPlannerTest, ColocatedJoinPushedToShards) { } } +TEST_F(DistributedPlannerTest, ColocatedLeftJoinEmitsLeftJoin) { + shard_map.add_table(TableShardConfig{ + "orders", "user_id", + {ShardInfo{"shard_1"}, ShardInfo{"shard_2"}, ShardInfo{"shard_3"}} + }); + + Parser parser; + const char* sql = "SELECT * FROM users LEFT JOIN orders ON users.id = orders.user_id"; + auto pr = parser.parse(sql, std::strlen(sql)); + ASSERT_EQ(pr.status, ParseResult::OK); + PlanBuilder builder(catalog, parser.arena()); + PlanNode* plan = builder.build(pr.ast); + DistributedPlanner dp(shard_map, catalog, parser.arena()); + PlanNode* dist = dp.distribute(plan); + ASSERT_NE(dist, nullptr); + + std::vector remotes; + find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes); + ASSERT_FALSE(remotes.empty()); + for (auto* rs : remotes) { + std::string remote(rs->remote_scan.remote_sql, rs->remote_scan.remote_sql_len); + EXPECT_NE(remote.find("LEFT JOIN"), std::string::npos) << remote; + } +} + +TEST_F(DistributedPlannerTest, SemiJoinSkipsLeftJoin) { + Parser parser; + const char* sql = "SELECT * FROM users LEFT JOIN orders ON users.id = orders.user_id"; + auto pr = parser.parse(sql, std::strlen(sql)); + ASSERT_EQ(pr.status, ParseResult::OK); + PlanBuilder builder(catalog, parser.arena()); + PlanNode* plan = builder.build(pr.ast); + DistributedPlanner dp(shard_map, catalog, parser.arena(), + &mock_executor, &functions); + PlanNode* dist = dp.distribute(plan); + ASSERT_NE(dist, nullptr); + std::vector remotes; + find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes); + for (auto* rs : remotes) { + std::string remote(rs->remote_scan.remote_sql, rs->remote_scan.remote_sql_len); + if (remote.find("users") != std::string::npos) + EXPECT_EQ(remote.find(" IN "), std::string::npos) << remote; + } +} + TEST_F(DistributedPlannerTest, CompositeColocatedJoinPushedToShards) { catalog.add_table("", "kv", { {"tenant_id", SqlType::make_int(), false}, @@ -1313,6 +1358,86 @@ TEST_F(DistributedPlannerTest, CompositeRangePrunesOnFirstKey) { EXPECT_EQ(count_remotes("SELECT * FROM users WHERE id BETWEEN 6 AND 10"), 1u); } +TEST_F(DistributedPlannerTest, ColocatedRightJoinEmitsRightJoin) { + shard_map.add_table(TableShardConfig{ + "orders", "user_id", + {ShardInfo{"shard_1"}, ShardInfo{"shard_2"}, ShardInfo{"shard_3"}} + }); + Parser parser; + const char* sql = "SELECT * FROM users RIGHT JOIN orders ON users.id = orders.user_id"; + auto pr = parser.parse(sql, std::strlen(sql)); + ASSERT_EQ(pr.status, ParseResult::OK); + PlanBuilder builder(catalog, parser.arena()); + PlanNode* plan = builder.build(pr.ast); + DistributedPlanner dp(shard_map, catalog, parser.arena()); + PlanNode* dist = dp.distribute(plan); + std::vector remotes; + find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes); + ASSERT_FALSE(remotes.empty()); + bool saw_right = false; + for (auto* rs : remotes) { + std::string remote(rs->remote_scan.remote_sql, rs->remote_scan.remote_sql_len); + if (remote.find("RIGHT JOIN") != std::string::npos) saw_right = true; + } + EXPECT_TRUE(saw_right); +} + +TEST_F(DistributedPlannerTest, OrEqualityPrunesUnionOfShards) { + TableShardConfig cfg; + cfg.table_name = "users"; + cfg.shard_key = "id"; + cfg.shards = {{"shard_1"}, {"shard_2"}, {"shard_3"}}; + cfg.strategy = RoutingStrategy::LIST; + cfg.list = {{true, 1, "", 0}, {true, 6, "", 1}, {true, 15, "", 2}}; + shard_map.add_table(cfg); + + Parser parser; + auto pr = parser.parse("SELECT * FROM users WHERE id = 1 OR id = 6", 42); + PlanBuilder builder(catalog, parser.arena()); + PlanNode* plan = builder.build(pr.ast); + DistributedPlanner dp(shard_map, catalog, parser.arena()); + PlanNode* dist = dp.distribute(plan); + std::vector remotes; + find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes); + EXPECT_EQ(remotes.size(), 2u); +} + +TEST_F(DistributedPlannerTest, CompositeHashPartialWhereScatters) { + TableShardConfig cfg; + cfg.table_name = "users"; + cfg.shard_key = "id+name"; + cfg.shards = {{"shard_1"}, {"shard_2"}, {"shard_3"}}; + shard_map.add_table(cfg); + + Parser parser; + auto pr = parser.parse("SELECT * FROM users WHERE id = 3", 32); + PlanBuilder builder(catalog, parser.arena()); + PlanNode* plan = builder.build(pr.ast); + DistributedPlanner dp(shard_map, catalog, parser.arena()); + PlanNode* dist = dp.distribute(plan); + std::vector remotes; + find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes); + EXPECT_EQ(remotes.size(), 3u); +} + +TEST_F(DistributedPlannerTest, CompositeHashBothKeysPrune) { + TableShardConfig cfg; + cfg.table_name = "users"; + cfg.shard_key = "id+age"; + cfg.shards = {{"shard_1"}, {"shard_2"}, {"shard_3"}}; + shard_map.add_table(cfg); + + Parser parser; + auto pr = parser.parse("SELECT * FROM users WHERE id = 3 AND age = 17", 45); + PlanBuilder builder(catalog, parser.arena()); + PlanNode* plan = builder.build(pr.ast); + DistributedPlanner dp(shard_map, catalog, parser.arena()); + PlanNode* dist = dp.distribute(plan); + std::vector remotes; + find_nodes(dist, PlanNodeType::REMOTE_SCAN, remotes); + EXPECT_EQ(remotes.size(), 1u); +} + TEST_F(DistributedPlannerTest, SemiJoinPrunesProbeShards) { Parser parser; const char* sql = "SELECT * FROM users JOIN orders ON users.id = orders.user_id"; diff --git a/tests/test_shard_battery.cpp b/tests/test_shard_battery.cpp new file mode 100644 index 0000000..e77b8d5 --- /dev/null +++ b/tests/test_shard_battery.cpp @@ -0,0 +1,781 @@ +#include +#include "sql_engine/shard_map.h" +#include "sql_engine/distributed_planner.h" +#include "sql_engine/plan_builder.h" +#include "sql_engine/dml_plan_builder.h" +#include "sql_engine/plan_executor.h" +#include "sql_engine/in_memory_catalog.h" +#include "sql_engine/function_registry.h" +#include "sql_engine/remote_executor.h" +#include "sql_engine/mutable_data_source.h" +#include "sql_engine/session.h" +#include "sql_engine/local_txn.h" +#include "sql_parser/parser.h" + +#include +#include +#include +#include +#include + +using namespace sql_engine; +using namespace sql_parser; + +namespace { + +StringRef sref(const char* s) { + return StringRef{s, static_cast(std::strlen(s))}; +} + +void find_nodes(PlanNode* n, PlanNodeType t, std::vector& out) { + if (!n) return; + if (n->type == t) out.push_back(n); + find_nodes(n->left, t, out); + find_nodes(n->right, t, out); + if (n->type == PlanNodeType::MERGE_SORT) { + for (uint16_t i = 0; i < n->merge_sort.child_count; ++i) + find_nodes(n->merge_sort.children[i], t, out); + } + if (n->type == PlanNodeType::MERGE_AGGREGATE) { + for (uint16_t i = 0; i < n->merge_aggregate.child_count; ++i) + find_nodes(n->merge_aggregate.children[i], t, out); + } +} + +TableShardConfig hash3() { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id"; + c.shards = {{"a"}, {"b"}, {"c"}}; + return c; +} + +} // namespace + +// ===================================================================== +// ShardMap routing battery +// ===================================================================== + +TEST(ShardBatteryHash, SameValueSameShard100) { + ShardMap map; + map.add_table(hash3()); + for (int i = -50; i < 50; ++i) { + size_t x = 0, y = 0; + ASSERT_TRUE(map.try_shard_index_for_int(sref("users"), i, x)); + ASSERT_TRUE(map.try_shard_index_for_int(sref("users"), i, y)); + EXPECT_EQ(x, y); + EXPECT_LT(x, 3u); + } +} + +TEST(ShardBatteryHash, StringIntMatchesIntWide) { + ShardMap map; + map.add_table(hash3()); + for (int i = -30; i < 30; ++i) { + std::string s = std::to_string(i); + size_t a = 0, b = 0; + ASSERT_TRUE(map.try_shard_index_for_int(sref("users"), i, a)); + ASSERT_TRUE(map.try_shard_index_for_string( + sref("users"), s.c_str(), static_cast(s.size()), b)); + EXPECT_EQ(a, b) << i; + } +} + +TEST(ShardBatteryHash, HitsAllThreeShards) { + ShardMap map; + map.add_table(hash3()); + bool seen[3] = {}; + for (int i = 0; i < 64; ++i) { + size_t x = 0; + ASSERT_TRUE(map.try_shard_index_for_int(sref("users"), i, x)); + seen[x] = true; + } + EXPECT_TRUE(seen[0] && seen[1] && seen[2]); +} + +TEST(ShardBatteryHash, UnknownTableFalse) { + ShardMap map; + size_t x = 7; + EXPECT_FALSE(map.try_shard_index_for_int(sref("nope"), 1, x)); + EXPECT_EQ(map.shard_index_for_int(sref("nope"), 1), static_cast(-1)); +} + +TEST(ShardBatteryHash, CaseInsensitiveTable) { + ShardMap map; + map.add_table(hash3()); + size_t a = 0, b = 0; + ASSERT_TRUE(map.try_shard_index_for_int(sref("USERS"), 5, a)); + ASSERT_TRUE(map.try_shard_index_for_int(sref("users"), 5, b)); + EXPECT_EQ(a, b); +} + +TEST(ShardBatteryHash, EmptyStringRoutes) { + ShardMap map; + map.add_table(hash3()); + size_t x = 9; + EXPECT_TRUE(map.try_shard_index_for_string(sref("users"), "", 0, x)); + EXPECT_LT(x, 3u); +} + +TEST(ShardBatteryHash, LongStringStable) { + ShardMap map; + map.add_table(hash3()); + const char* s = "the-quick-brown-fox-jumps-over-the-lazy-dog"; + uint32_t n = static_cast(std::strlen(s)); + size_t a = 0, b = 0; + ASSERT_TRUE(map.try_shard_index_for_string(sref("users"), s, n, a)); + ASSERT_TRUE(map.try_shard_index_for_string(sref("users"), s, n, b)); + EXPECT_EQ(a, b); +} + +TEST(ShardBatteryRange, Bounds) { + TableShardConfig c = hash3(); + c.strategy = RoutingStrategy::RANGE; + c.ranges = {{0, 0}, {10, 1}, {100, 2}}; + ShardMap map; + map.add_table(c); + size_t x = 9; + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), -1, x)); + EXPECT_EQ(x, 0u); + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 0, x)); + EXPECT_EQ(x, 0u); + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 1, x)); + EXPECT_EQ(x, 1u); + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 10, x)); + EXPECT_EQ(x, 1u); + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 11, x)); + EXPECT_EQ(x, 2u); + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 9999, x)); + EXPECT_EQ(x, 2u); +} + +TEST(ShardBatteryRange, CollectWindow) { + TableShardConfig c = hash3(); + c.strategy = RoutingStrategy::RANGE; + c.ranges = {{5, 0}, {10, 1}, {100, 2}}; + ShardMap map; + map.add_table(c); + std::vector out; + map.collect_int_range_shards(sref("users"), 6, 10, out); + ASSERT_EQ(out.size(), 1u); + EXPECT_EQ(out[0], 1u); +} + +TEST(ShardBatteryRange, StringUnroutable) { + TableShardConfig c = hash3(); + c.strategy = RoutingStrategy::RANGE; + c.ranges = {{5, 0}, {10, 1}}; + ShardMap map; + map.add_table(c); + size_t x = 0; + EXPECT_FALSE(map.try_shard_index_for_string(sref("users"), "3", 1, x)); +} + +TEST(ShardBatteryList, MappedAndMiss) { + TableShardConfig c = hash3(); + c.strategy = RoutingStrategy::LIST; + c.list = {{true, 1, "", 0}, {true, 2, "", 0}, {true, 9, "", 2}}; + ShardMap map; + map.add_table(c); + size_t x = 9; + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 1, x)); + EXPECT_EQ(x, 0u); + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 9, x)); + EXPECT_EQ(x, 2u); + EXPECT_FALSE(map.try_shard_index_for_int(sref("users"), 3, x)); + EXPECT_EQ(map.shard_index_for_int(sref("users"), 3), static_cast(-1)); +} + +TEST(ShardBatteryList, BetweenTwoShards) { + TableShardConfig c = hash3(); + c.strategy = RoutingStrategy::LIST; + c.list = {{true, 1, "", 0}, {true, 5, "", 1}, {true, 9, "", 2}}; + ShardMap map; + map.add_table(c); + std::vector out; + map.collect_int_list_shards(sref("users"), 1, 5, out); + ASSERT_EQ(out.size(), 2u); +} + +TEST(ShardBatteryList, StringKeys) { + TableShardConfig c = hash3(); + c.strategy = RoutingStrategy::LIST; + c.list = {{false, 0, "east", 0}, {false, 0, "west", 1}}; + ShardMap map; + map.add_table(c); + size_t x = 9; + EXPECT_TRUE(map.try_shard_index_for_string(sref("users"), "east", 4, x)); + EXPECT_EQ(x, 0u); + EXPECT_FALSE(map.try_shard_index_for_string(sref("users"), "north", 5, x)); +} + +TEST(ShardBatteryComposite, HashTwoPartsStable) { + TableShardConfig c; + c.table_name = "kv"; + c.shard_key = "t+id"; + c.shards = {{"a"}, {"b"}, {"c"}}; + ShardMap map; + map.add_table(c); + ShardKeyPart p[] = {{true, 1, nullptr, 0}, {true, 2, nullptr, 0}}; + size_t a = 0, b = 0; + ASSERT_TRUE(map.try_shard_index_for_parts(sref("kv"), p, 2, a)); + ASSERT_TRUE(map.try_shard_index_for_parts(sref("kv"), p, 2, b)); + EXPECT_EQ(a, b); +} + +TEST(ShardBatteryComposite, RangeUsesFirstOnly) { + TableShardConfig c; + c.table_name = "kv"; + c.shard_key = "t+id"; + c.shards = {{"a"}, {"b"}}; + c.strategy = RoutingStrategy::RANGE; + c.ranges = {{5, 0}, {100, 1}}; + ShardMap map; + map.add_table(c); + ShardKeyPart lo[] = {{true, 1, nullptr, 0}, {true, 99, nullptr, 0}}; + ShardKeyPart hi[] = {{true, 9, nullptr, 0}, {true, 0, nullptr, 0}}; + size_t a = 9, b = 9; + ASSERT_TRUE(map.try_shard_index_for_parts(sref("kv"), lo, 2, a)); + ASSERT_TRUE(map.try_shard_index_for_parts(sref("kv"), hi, 2, b)); + EXPECT_EQ(a, 0u); + EXPECT_EQ(b, 1u); +} + +TEST(ShardBatteryComposite, GetKeysSplit) { + TableShardConfig c; + c.table_name = "kv"; + c.shard_key = "tenant_id+user_id"; + c.shards = {{"a"}, {"b"}}; + ShardMap map; + map.add_table(c); + const auto& keys = map.get_shard_keys(sref("kv")); + ASSERT_EQ(keys.size(), 2u); + EXPECT_EQ(keys[0], "tenant_id"); + EXPECT_EQ(keys[1], "user_id"); +} + +TEST(ShardBatterySameRouting, IdenticalLayouts) { + ShardMap map; + TableShardConfig a = hash3(); + TableShardConfig b = hash3(); + b.table_name = "orders"; + b.shard_key = "user_id"; + map.add_table(a); + map.add_table(b); + EXPECT_TRUE(map.same_routing(sref("users"), sref("orders"))); +} + +TEST(ShardBatterySameRouting, DifferentShardCount) { + ShardMap map; + map.add_table(hash3()); + TableShardConfig b; + b.table_name = "orders"; + b.shard_key = "id"; + b.shards = {{"a"}, {"b"}}; + map.add_table(b); + EXPECT_FALSE(map.same_routing(sref("users"), sref("orders"))); +} + +// ===================================================================== +// Planner prune battery +// ===================================================================== + +class PruneBattery : public ::testing::Test { +protected: + InMemoryCatalog catalog; + ShardMap shards; + const TableInfo* users = nullptr; + + void SetUp() override { + catalog.add_table("", "users", { + {"id", SqlType::make_int(), false}, + {"name", SqlType::make_varchar(255), true}, + {"age", SqlType::make_int(), true}, + }); + users = catalog.get_table(sref("users")); + shards.add_table({"users", "id", {{"s0"}, {"s1"}, {"s2"}}}); + } + + size_t remotes(const char* sql) { + Parser p; + auto pr = p.parse(sql, std::strlen(sql)); + if (pr.status != ParseResult::OK) return 999; + PlanBuilder b(catalog, p.arena()); + PlanNode* plan = b.build(pr.ast); + DistributedPlanner dp(shards, catalog, p.arena()); + PlanNode* dist = dp.distribute(plan); + if (!dist && dp.last_error()) return 0; + std::vector out; + find_nodes(dist, PlanNodeType::REMOTE_SCAN, out); + return out.size(); + } + + void set_range() { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + c.strategy = RoutingStrategy::RANGE; + c.ranges = {{5, 0}, {10, 1}, {100000, 2}}; + shards.add_table(c); + } + + void set_list() { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + c.strategy = RoutingStrategy::LIST; + c.list = {{true, 1, "", 0}, {true, 6, "", 1}, {true, 7, "", 1}, + {true, 15, "", 2}}; + shards.add_table(c); + } +}; + +TEST_F(PruneBattery, HashPointIsSingle) { EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 4"), 1u); } +TEST_F(PruneBattery, HashNoWhereIsThree) { EXPECT_EQ(remotes("SELECT * FROM users"), 3u); } +TEST_F(PruneBattery, HashNonKeyIsThree) { EXPECT_EQ(remotes("SELECT * FROM users WHERE age = 1"), 3u); } +TEST_F(PruneBattery, HashInTwoValues) { + size_t n = remotes("SELECT * FROM users WHERE id IN (1, 2)"); + EXPECT_GE(n, 1u); + EXPECT_LE(n, 2u); +} +TEST_F(PruneBattery, HashPlaceholderScatters) { EXPECT_EQ(remotes("SELECT * FROM users WHERE id = ?"), 3u); } +TEST_F(PruneBattery, RangeLe5) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id <= 5"), 1u); } +TEST_F(PruneBattery, RangeGt10) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id > 10"), 1u); } +TEST_F(PruneBattery, RangeBetween6And10) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id BETWEEN 6 AND 10"), 1u); } +TEST_F(PruneBattery, RangeGe1Le15) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id >= 1 AND id <= 15"), 3u); } +TEST_F(PruneBattery, RangeLt1) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id < 1"), 1u); } +TEST_F(PruneBattery, RangeEq3) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 3"), 1u); } +TEST_F(PruneBattery, RangeEq7) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 7"), 1u); } +TEST_F(PruneBattery, RangeEq20) { set_range(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 20"), 1u); } +TEST_F(PruneBattery, ListEq1) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 1"), 1u); } +TEST_F(PruneBattery, ListEq6) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 6"), 1u); } +TEST_F(PruneBattery, ListMissScatters) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 99"), 3u); } +TEST_F(PruneBattery, ListBetween67) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id BETWEEN 6 AND 7"), 1u); } +TEST_F(PruneBattery, ListBetweenAll) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id BETWEEN 1 AND 15"), 3u); } +TEST_F(PruneBattery, ListInMappedAndMiss) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id IN (6, 99)"), 1u); } +TEST_F(PruneBattery, ListInTwoShards) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id IN (6, 15)"), 2u); } +TEST_F(PruneBattery, ListOrTwoShards) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 1 OR id = 6"), 2u); } +TEST_F(PruneBattery, ListEmptyBetweenScatters) { set_list(); EXPECT_EQ(remotes("SELECT * FROM users WHERE id BETWEEN 2 AND 5"), 3u); } +TEST_F(PruneBattery, AndWithNonKeyStillPrunes) { + EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 4 AND age > 0"), 1u); +} +TEST_F(PruneBattery, SelectStarCount) { EXPECT_GE(remotes("SELECT COUNT(*) FROM users"), 3u); } +TEST_F(PruneBattery, SelectStarCountPoint) { EXPECT_GE(remotes("SELECT COUNT(*) FROM users WHERE id = 4"), 1u); } +TEST_F(PruneBattery, OrderByScatter) { EXPECT_GE(remotes("SELECT * FROM users ORDER BY name"), 3u); } +TEST_F(PruneBattery, LimitScatter) { EXPECT_EQ(remotes("SELECT * FROM users LIMIT 5"), 3u); } +TEST_F(PruneBattery, DistinctScatter) { EXPECT_EQ(remotes("SELECT DISTINCT name FROM users"), 3u); } + +TEST_F(PruneBattery, UnknownTableErrors) { + catalog.add_table("", "ghost", {{"id", SqlType::make_int(), false}}); + Parser p; + auto pr = p.parse("SELECT * FROM ghost", 19); + PlanBuilder b(catalog, p.arena()); + PlanNode* plan = b.build(pr.ast); + DistributedPlanner dp(shards, catalog, p.arena()); + EXPECT_EQ(dp.distribute(plan), nullptr); + ASSERT_NE(dp.last_error(), nullptr); +} + +TEST_F(PruneBattery, CompositePartialScatters) { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id+age"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + shards.add_table(c); + EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 3"), 3u); +} + +TEST_F(PruneBattery, CompositeBothKeysPrune) { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id+age"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + shards.add_table(c); + EXPECT_EQ(remotes("SELECT * FROM users WHERE id = 3 AND age = 17"), 1u); +} + +TEST_F(PruneBattery, CompositeRangeFirstKey) { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id+name"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + c.strategy = RoutingStrategy::RANGE; + c.ranges = {{5, 0}, {10, 1}, {100000, 2}}; + shards.add_table(c); + EXPECT_EQ(remotes("SELECT * FROM users WHERE id <= 5"), 1u); + EXPECT_EQ(remotes("SELECT * FROM users WHERE id BETWEEN 6 AND 10"), 1u); +} + +// ===================================================================== +// DML battery (compact mock) +// ===================================================================== + +struct BatBackend { + InMemoryCatalog catalog; + FunctionRegistry functions; + Arena arena{65536, 1048576}; + std::map srcs; + std::vector> owned; + BatBackend() { functions.register_builtins(); } + void add_table(const char* n, std::initializer_list cols) { + catalog.add_table("", n, cols); + const TableInfo* ti = catalog.get_table(sref(n)); + auto* s = new InMemoryMutableDataSource(ti, arena); + srcs[n] = s; + owned.emplace_back(s); + } +}; + +class BatExec : public RemoteExecutor { +public: + void add_backend(const std::string& n) { b_[n] = std::make_unique(); } + BatBackend* get(const std::string& n) { + auto it = b_.find(n); + return it == b_.end() ? nullptr : it->second.get(); + } + void add_users_all() { + for (auto& kv : b_) { + kv.second->add_table("users", { + {"id", SqlType::make_int(), false}, + {"name", SqlType::make_varchar(255), true}, + {"age", SqlType::make_int(), true}, + }); + } + } + ResultSet execute(const char* name, StringRef sql) override { + auto* bd = get(name); + if (!bd) return ResultSet::fail("unknown backend"); + std::string s(sql.ptr, sql.len); + if (s.size() > 11 && s.compare(s.size() - 11, 11, " FOR UPDATE") == 0) + s.resize(s.size() - 11); + if (s.rfind("INSERT", 0) == 0 || s.rfind("UPDATE", 0) == 0 || s.rfind("DELETE", 0) == 0) { + run_dml(bd, s); + return {}; + } + Parser p; + auto pr = p.parse(s.c_str(), s.size()); + if (pr.status != ParseResult::OK || !pr.ast) return {}; + PlanBuilder b(bd->catalog, p.arena()); + PlanNode* plan = b.build(pr.ast); + if (!plan) return {}; + PlanExecutor ex(bd->functions, bd->catalog, p.arena()); + for (auto& t : bd->srcs) ex.add_mutable_data_source(t.first.c_str(), t.second); + ResultSet rs = ex.execute(plan); + rs.ok = true; + return rs; + } + DmlResult execute_dml(const char* name, StringRef sql) override { + auto* bd = get(name); + if (!bd) { + DmlResult r; + r.error_message = "unknown backend"; + return r; + } + return run_dml(bd, std::string(sql.ptr, sql.len)); + } + size_t total(const char* table) { + size_t n = 0; + for (auto& kv : b_) { + auto it = kv.second->srcs.find(table); + if (it != kv.second->srcs.end()) n += it->second->row_count(); + } + return n; + } + size_t on(const char* backend, const char* table) { + auto* bd = get(backend); + if (!bd) return 0; + auto it = bd->srcs.find(table); + return it == bd->srcs.end() ? 0 : it->second->row_count(); + } +private: + std::map> b_; + DmlResult run_dml(BatBackend* bd, const std::string& sql) { + Parser p; + auto pr = p.parse(sql.c_str(), sql.size()); + if (pr.status != ParseResult::OK || !pr.ast) { + DmlResult r; + r.error_message = "parse error"; + return r; + } + DmlPlanBuilder b(bd->catalog, p.arena()); + PlanNode* plan = b.build(pr.ast); + if (!plan) { + DmlResult r; + r.error_message = "plan error"; + return r; + } + PlanExecutor ex(bd->functions, bd->catalog, p.arena()); + for (auto& t : bd->srcs) ex.add_mutable_data_source(t.first.c_str(), t.second); + return ex.execute_dml(plan); + } +}; + +class DmlBattery : public ::testing::Test { +protected: + Arena arena{65536, 1048576}; + InMemoryCatalog catalog; + FunctionRegistry functions; + ShardMap shards; + BatExec exec; + + void SetUp() override { + functions.register_builtins(); + catalog.add_table("", "users", { + {"id", SqlType::make_int(), false}, + {"name", SqlType::make_varchar(255), true}, + {"age", SqlType::make_int(), true}, + }); + shards.add_table({"users", "id", {{"s0"}, {"s1"}, {"s2"}}}); + exec.add_backend("s0"); + exec.add_backend("s1"); + exec.add_backend("s2"); + exec.add_users_all(); + } + + const char* backend(int64_t id) { + return shards.get_shards(sref("users"))[shards.shard_index_for_int(sref("users"), id)] + .backend_name.c_str(); + } + + DmlResult dml(const char* sql) { + Parser p; + auto pr = p.parse(sql, std::strlen(sql)); + if (pr.status != ParseResult::OK || !pr.ast) { + DmlResult r; + r.error_message = "parse"; + return r; + } + DmlPlanBuilder b(catalog, p.arena()); + PlanNode* plan = b.build(pr.ast); + DistributedPlanner dp(shards, catalog, p.arena(), &exec, &functions); + PlanNode* dist = dp.distribute_dml(plan); + if (dp.last_error()) { + DmlResult r; + r.error_message = dp.last_error(); + return r; + } + DmlResult total; + total.success = true; + std::function walk = [&](PlanNode* n) { + if (!n) return; + if (n->type == PlanNodeType::REMOTE_SCAN) { + StringRef s{n->remote_scan.remote_sql, n->remote_scan.remote_sql_len}; + DmlResult r = exec.execute_dml(n->remote_scan.backend_name, s); + if (!r.success) { + total.success = false; + total.error_message = r.error_message; + } + total.affected_rows += r.affected_rows; + return; + } + walk(n->left); + walk(n->right); + }; + walk(dist); + return total; + } + + ResultSet sel(const char* sql) { + Parser p; + auto pr = p.parse(sql, std::strlen(sql)); + PlanBuilder b(catalog, p.arena()); + PlanNode* plan = b.build(pr.ast); + DistributedPlanner dp(shards, catalog, p.arena(), &exec, &functions); + PlanNode* dist = dp.distribute(plan); + PlanExecutor ex(functions, catalog, p.arena()); + ex.set_remote_executor(&exec); + return ex.execute(dist); + } +}; + +TEST_F(DmlBattery, InsertThenSelect0) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (0, 'n', 1)").success); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 0").row_count(), 1u); +} +TEST_F(DmlBattery, InsertThenSelect1) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (1, 'n', 1)").success); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 1").row_count(), 1u); +} +TEST_F(DmlBattery, InsertThenSelect2) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (2, 'n', 1)").success); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 2").row_count(), 1u); +} +TEST_F(DmlBattery, InsertThenSelect7) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (7, 'n', 1)").success); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 7").row_count(), 1u); +} +TEST_F(DmlBattery, InsertThenSelect13) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (13, 'n', 1)").success); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 13").row_count(), 1u); +} +TEST_F(DmlBattery, InsertThenSelectNeg3) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (-3, 'n', 1)").success); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = -3").row_count(), 1u); +} +TEST_F(DmlBattery, InsertStringThenIntSelect) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES ('5', 'n', 1)").success); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 5").row_count(), 1u); +} +TEST_F(DmlBattery, InsertLandsOnHashedShard) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (8, 'n', 1)").success); + EXPECT_EQ(exec.on(backend(8), "users"), 1u); + EXPECT_EQ(exec.total("users"), 1u); +} +TEST_F(DmlBattery, MissingKeyErrors) { + EXPECT_FALSE(dml("INSERT INTO users (name, age) VALUES ('n', 1)").success); +} +TEST_F(DmlBattery, NonLiteralErrors) { + EXPECT_FALSE(dml("INSERT INTO users (id, name, age) VALUES (1+1, 'n', 1)").success); +} +TEST_F(DmlBattery, DeletePoint) { + dml("INSERT INTO users (id, name, age) VALUES (4, 'n', 1)"); + EXPECT_TRUE(dml("DELETE FROM users WHERE id = 4").success); + EXPECT_EQ(exec.total("users"), 0u); +} +TEST_F(DmlBattery, DeleteScatterNone) { + dml("INSERT INTO users (id, name, age) VALUES (4, 'n', 1)"); + EXPECT_TRUE(dml("DELETE FROM users WHERE age = 99").success); + EXPECT_EQ(exec.total("users"), 1u); +} +TEST_F(DmlBattery, UpdateNonKeyPoint) { + dml("INSERT INTO users (id, name, age) VALUES (4, 'n', 1)"); + EXPECT_TRUE(dml("UPDATE users SET age = 9 WHERE id = 4").success); + auto rs = sel("SELECT age FROM users WHERE id = 4"); + ASSERT_EQ(rs.row_count(), 1u); + EXPECT_EQ(rs.rows[0].get(0).int_val, 9); +} +TEST_F(DmlBattery, UpdateMoveCrossShard) { + dml("INSERT INTO users (id, name, age) VALUES (3, 'n', 1)"); + int64_t dest = 3; + for (int64_t i = 4; i < 80; ++i) { + if (std::strcmp(backend(i), backend(3)) != 0) { dest = i; break; } + } + ASSERT_NE(dest, 3); + std::string sql = "UPDATE users SET id = " + std::to_string(dest) + " WHERE id = 3"; + EXPECT_TRUE(dml(sql.c_str()).success); + EXPECT_EQ(exec.on(backend(3), "users"), 0u); + EXPECT_EQ(exec.on(backend(dest), "users"), 1u); +} +TEST_F(DmlBattery, UpdateQualifiedMove) { + dml("INSERT INTO users (id, name, age) VALUES (3, 'n', 1)"); + int64_t dest = 9; + if (std::strcmp(backend(3), backend(9)) == 0) dest = 8; + std::string sql = "UPDATE users SET users.id = " + std::to_string(dest) + " WHERE id = 3"; + EXPECT_TRUE(dml(sql.c_str()).success); + EXPECT_EQ(sel(("SELECT name FROM users WHERE id = " + std::to_string(dest)).c_str()).row_count(), 1u); +} +TEST_F(DmlBattery, UpdateMoveNoRow) { + dml("INSERT INTO users (id, name, age) VALUES (3, 'n', 1)"); + EXPECT_TRUE(dml("UPDATE users SET id = 9 WHERE id = 99").success); + EXPECT_EQ(exec.total("users"), 1u); +} +TEST_F(DmlBattery, MultiInsertThenEachSelect) { + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (10, 'a', 1), (11, 'b', 1), (12, 'c', 1)").success); + EXPECT_EQ(exec.total("users"), 3u); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 10").row_count(), 1u); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 11").row_count(), 1u); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 12").row_count(), 1u); +} +TEST_F(DmlBattery, SessionUnknownTableErrors) { + catalog.add_table("", "ghost", {{"id", SqlType::make_int(), false}}); + LocalTransactionManager txn(arena); + Session session(catalog, txn); + session.set_remote_executor(&exec); + session.set_shard_map(&shards); + auto rs = session.execute_query("SELECT * FROM ghost"); + EXPECT_FALSE(rs.ok); +} +TEST_F(DmlBattery, SessionCountTwice) { + dml("INSERT INTO users (id, name, age) VALUES (1, 'a', 2)"); + dml("INSERT INTO users (id, name, age) VALUES (2, 'b', 3)"); + LocalTransactionManager txn(arena); + Session session(catalog, txn); + session.set_remote_executor(&exec); + session.set_shard_map(&shards); + auto a = session.execute_query("SELECT COUNT(*) FROM users"); + auto b = session.execute_query("SELECT COUNT(*) FROM users"); + ASSERT_EQ(a.row_count(), 1u); + ASSERT_EQ(b.row_count(), 1u); + EXPECT_EQ(a.rows[0].get(0).to_int64(), 2); + EXPECT_EQ(b.rows[0].get(0).to_int64(), 2); +} +TEST_F(DmlBattery, RangeInsertSelect) { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + c.strategy = RoutingStrategy::RANGE; + c.ranges = {{5, 0}, {10, 1}, {1000, 2}}; + shards.add_table(c); + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (3, 'lo', 1)").success); + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (7, 'mid', 1)").success); + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (20, 'hi', 1)").success); + EXPECT_EQ(exec.on("s0", "users"), 1u); + EXPECT_EQ(exec.on("s1", "users"), 1u); + EXPECT_EQ(exec.on("s2", "users"), 1u); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 3").row_count(), 1u); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 7").row_count(), 1u); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 20").row_count(), 1u); +} +TEST_F(DmlBattery, ListInsertUnmappedErrors) { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + c.strategy = RoutingStrategy::LIST; + c.list = {{true, 1, "", 0}, {true, 2, "", 1}}; + shards.add_table(c); + EXPECT_FALSE(dml("INSERT INTO users (id, name, age) VALUES (99, 'x', 1)").success); + EXPECT_EQ(exec.total("users"), 0u); +} +TEST_F(DmlBattery, ListInsertMapped) { + TableShardConfig c; + c.table_name = "users"; + c.shard_key = "id"; + c.shards = {{"s0"}, {"s1"}, {"s2"}}; + c.strategy = RoutingStrategy::LIST; + c.list = {{true, 1, "", 0}, {true, 2, "", 1}}; + shards.add_table(c); + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (1, 'x', 1)").success); + EXPECT_EQ(exec.on("s0", "users"), 1u); + EXPECT_EQ(sel("SELECT name FROM users WHERE id = 1").row_count(), 1u); +} + +#define HASH_ROUNDTRIP(n) \ +TEST_F(DmlBattery, InsertSelect_##n) { \ + EXPECT_TRUE(dml("INSERT INTO users (id, name, age) VALUES (" #n ", 'x', 1)").success); \ + EXPECT_EQ(sel("SELECT id FROM users WHERE id = " #n).row_count(), 1u); \ +} + +HASH_ROUNDTRIP(20) +HASH_ROUNDTRIP(21) +HASH_ROUNDTRIP(22) +HASH_ROUNDTRIP(23) +HASH_ROUNDTRIP(24) +HASH_ROUNDTRIP(25) +HASH_ROUNDTRIP(26) +HASH_ROUNDTRIP(27) +HASH_ROUNDTRIP(28) +HASH_ROUNDTRIP(29) +HASH_ROUNDTRIP(30) +HASH_ROUNDTRIP(31) +HASH_ROUNDTRIP(32) +HASH_ROUNDTRIP(33) +HASH_ROUNDTRIP(34) +HASH_ROUNDTRIP(35) +HASH_ROUNDTRIP(36) +HASH_ROUNDTRIP(37) +HASH_ROUNDTRIP(38) +HASH_ROUNDTRIP(39) +HASH_ROUNDTRIP(40) +HASH_ROUNDTRIP(41) +HASH_ROUNDTRIP(42) +HASH_ROUNDTRIP(43) +HASH_ROUNDTRIP(44) +HASH_ROUNDTRIP(45) +HASH_ROUNDTRIP(46) +HASH_ROUNDTRIP(47) +HASH_ROUNDTRIP(48) +HASH_ROUNDTRIP(49) diff --git a/tests/test_shard_map.cpp b/tests/test_shard_map.cpp index a656e82..c985c06 100644 --- a/tests/test_shard_map.cpp +++ b/tests/test_shard_map.cpp @@ -10,6 +10,7 @@ #include #include +#include using namespace sql_engine; using sql_parser::StringRef; @@ -47,6 +48,75 @@ TEST(ShardMapHashTest, IsDeterministic) { } } +TEST(ShardMapHashTest, LeadingZerosRouteLikeInt) { + ShardMap map; + map.add_table(make_two_shards(RoutingStrategy::HASH)); + size_t a = map.shard_index_for_int(sref("users"), 3); + size_t b = 99; + ASSERT_TRUE(map.try_shard_index_for_string(sref("users"), "03", 2, b)); + EXPECT_EQ(a, b); +} + +TEST(ShardMapHashTest, NonNumericStringDoesNotCoerce) { + TableShardConfig cfg; + cfg.table_name = "users"; + cfg.shard_key = "id"; + cfg.shards = {ShardInfo{"a"}, ShardInfo{"b"}, ShardInfo{"c"}, ShardInfo{"d"}}; + ShardMap map; + map.add_table(cfg); + size_t as_int = map.shard_index_for_int(sref("users"), 3); + bool differs = false; + const char* samples[] = {"3x", "3.0", " 3", "3 "}; + for (const char* s : samples) { + size_t idx = 99; + ASSERT_TRUE(map.try_shard_index_for_string( + sref("users"), s, static_cast(std::strlen(s)), idx)); + if (idx != as_int) differs = true; + } + EXPECT_TRUE(differs); +} + +TEST(ShardMapHashTest, EmptyStringIsUnparsed) { + ShardMap map; + map.add_table(make_two_shards(RoutingStrategy::HASH)); + size_t idx = 99; + EXPECT_TRUE(map.try_shard_index_for_string(sref("users"), "", 0, idx)); + EXPECT_LT(idx, 2u); +} + +TEST(ShardMapListTest, StringIntegerDoesNotCoerceToIntKey) { + TableShardConfig cfg = make_two_shards(RoutingStrategy::LIST); + cfg.list = {ShardListEntry{true, 3, "", 1}}; + ShardMap map; + map.add_table(cfg); + size_t idx = 99; + EXPECT_FALSE(map.try_shard_index_for_string(sref("users"), "3", 1, idx)); + EXPECT_TRUE(map.try_shard_index_for_int(sref("users"), 3, idx)); + EXPECT_EQ(idx, 1u); +} + +TEST(ShardMapRangeTest, StringIntegerIsUnroutable) { + TableShardConfig cfg = make_two_shards(RoutingStrategy::RANGE); + cfg.ranges = {ShardRange{5, 0}, ShardRange{100, 1}}; + ShardMap map; + map.add_table(cfg); + size_t idx = 99; + EXPECT_FALSE(map.try_shard_index_for_string(sref("users"), "3", 1, idx)); +} + +TEST(ShardMapHashTest, StringIntegerRoutesLikeInt) { + ShardMap map; + map.add_table(make_two_shards(RoutingStrategy::HASH)); + for (int i = -20; i < 20; ++i) { + std::string s = std::to_string(i); + size_t a = map.shard_index_for_int(sref("users"), i); + size_t b = 99; + ASSERT_TRUE(map.try_shard_index_for_string( + sref("users"), s.c_str(), static_cast(s.size()), b)); + EXPECT_EQ(a, b) << "for " << i; + } +} + TEST(ShardMapHashTest, DistributesAcrossBothShards) { // FNV-1a with modulo 2 should hit both shards for the small range // 0..63 — the test is loose because we don't want to over-specify @@ -155,7 +225,8 @@ TEST(ShardMapRangeTest, StringKeyDoesNotRouteSpuriously) { ShardMap map; map.add_table(cfg); - EXPECT_EQ(map.shard_index_for_string(sref("users"), "anything", 8), 0u); + EXPECT_EQ(map.shard_index_for_string(sref("users"), "anything", 8), + static_cast(-1)); } // ---------------------------------------------------------------------- diff --git a/tools/engine_stress_test.cpp b/tools/engine_stress_test.cpp index 492df7c..47f1d07 100644 --- a/tools/engine_stress_test.cpp +++ b/tools/engine_stress_test.cpp @@ -385,7 +385,8 @@ static const std::string& backend_for_int_key( { StringRef t{table, static_cast(std::strlen(table))}; if (!map.has_table(t) || map.get_shards(t).empty()) return fallback.front(); - size_t idx = map.shard_index_for_int(t, key); + size_t idx = 0; + if (!map.try_shard_index_for_int(t, key, idx)) return fallback.front(); return map.get_shards(t)[idx].backend_name; }