Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
162 changes: 161 additions & 1 deletion include/sql_engine/distributed_planner.h
Original file line number Diff line number Diff line change
Expand Up @@ -357,7 +357,13 @@ class DistributedPlanner {
if (keys.empty()) return all_shards;

std::vector<size_t> target_indices;
if (keys.size() > 1) {
if (keys.size() > 1 &&
shards_.routing_strategy(table->table_name) == RoutingStrategy::RANGE) {
sql_parser::StringRef first{keys[0].c_str(),
static_cast<uint32_t>(keys[0].size())};
extract_shard_targets(where_expr, first, table->table_name,
all_shards.size(), target_indices);
} else if (keys.size() > 1) {
extract_composite_targets(where_expr, keys, table->table_name, target_indices);
} else {
sql_parser::StringRef shard_key{keys[0].c_str(),
Expand Down Expand Up @@ -1203,6 +1209,157 @@ class DistributedPlanner {
return current ? current : join_node;
}

bool column_on_table(const sql_parser::AstNode* node, const TableInfo* table,
const TableInfo* other) const {
if (!node || !table) return false;
if (node->type == sql_parser::NodeType::NODE_QUALIFIED_NAME) {
const sql_parser::AstNode* t = node->first_child;
if (!t) return false;
sql_parser::StringRef tn = t->value();
if (table->table_name.equals_ci(tn.ptr, tn.len)) return true;
if (table->alias.ptr && table->alias.equals_ci(tn.ptr, tn.len)) return true;
return false;
}
if (node->type == sql_parser::NodeType::NODE_COLUMN_REF ||
node->type == sql_parser::NodeType::NODE_IDENTIFIER) {
if (!catalog_.get_column(table, node->value())) return false;
if (other && catalog_.get_column(other, node->value())) return false;
return true;
}
return false;
}

const sql_parser::AstNode* probe_key_in_join(const sql_parser::AstNode* cond,
const TableInfo* probe,
const TableInfo* build) const {
if (!cond || !probe || !build) return nullptr;
const auto& keys = shards_.get_shard_keys(probe->table_name);
if (keys.size() != 1) return nullptr;
sql_parser::StringRef sk{keys[0].c_str(), static_cast<uint32_t>(keys[0].size())};
std::vector<std::pair<const sql_parser::AstNode*, const sql_parser::AstNode*>> eqs;
collect_eq_pairs(cond, eqs);
for (const auto& eq : eqs) {
if (is_shard_key_ref(eq.first, sk) && column_on_table(eq.second, build, probe))
return eq.first;
if (is_shard_key_ref(eq.second, sk) && column_on_table(eq.first, build, probe))
return eq.second;
}
return nullptr;
}

sql_parser::AstNode* make_in_list_on_column(const sql_parser::AstNode* col,
const std::vector<Value>& values) {
if (!col || values.empty()) return nullptr;
sql_parser::AstNode* stub = sql_parser::make_node(
arena_, sql_parser::NodeType::NODE_IN_LIST,
sql_parser::StringRef{nullptr, 0});
sql_parser::AstNode* col_copy = sql_parser::make_node(
arena_, col->type, col->value(), col->flags);
col_copy->first_child = col->first_child;
stub->add_child(col_copy);
return build_in_list_from_values(stub, values);
}

sql_parser::AstNode* and_preds(const sql_parser::AstNode* a,
const sql_parser::AstNode* b) {
if (!a) return const_cast<sql_parser::AstNode*>(b);
if (!b) return const_cast<sql_parser::AstNode*>(a);
sql_parser::AstNode* n = sql_parser::make_node(
arena_, sql_parser::NodeType::NODE_BINARY_OP,
sql_parser::StringRef{"AND", 3});
n->add_child(const_cast<sql_parser::AstNode*>(a));
n->add_child(const_cast<sql_parser::AstNode*>(b));
return n;
}

std::vector<Value> collect_build_join_keys(const TableInfo* build,
const sql_parser::AstNode* where_expr,
const sql_parser::AstNode* join_eq_other) {
std::vector<Value> out;
if (!build || !join_eq_other || !remote_executor_) return out;
const sql_parser::AstNode* proj[1] = {join_eq_other};
const auto& shards = shards_.get_shards(build->table_name);
if (shards.empty()) return out;
std::vector<ShardInfo> targets = shards;
if (shards_.is_sharded(build->table_name) && shards.size() > 1)
return out;
sql_parser::StringRef sql = qb_.build_select(
build, where_expr, proj, 1, nullptr, 0,
nullptr, nullptr, 0, -1, true);
ResultSet rs = remote_executor_->execute(shards[0].backend_name.c_str(), sql);
for (const auto& row : rs.rows) {
if (row.column_count > 0 && value_is_routable(row.get(0)))
out.push_back(copy_value_arena(row.get(0)));
}
return out;
}

const sql_parser::AstNode* other_eq_side(const sql_parser::AstNode* cond,
const sql_parser::AstNode* probe_key) const {
std::vector<std::pair<const sql_parser::AstNode*, const sql_parser::AstNode*>> eqs;
collect_eq_pairs(cond, eqs);
for (const auto& eq : eqs) {
if (eq.first == probe_key) return eq.second;
if (eq.second == probe_key) return eq.first;
}
return nullptr;
}

PlanNode* try_semijoin_prune(PlanNode* join_node,
const TableInfo* left_table,
const TableInfo* right_table) {
if (!join_node || !remote_executor_ || !join_node->join.condition)
return nullptr;
if (!left_table || !right_table) return nullptr;

bool ls = shards_.is_sharded(left_table->table_name);
bool rs = shards_.is_sharded(right_table->table_name);
if (ls == rs) return nullptr;

const TableInfo* probe = ls ? left_table : right_table;
const TableInfo* build = ls ? right_table : left_table;
bool probe_is_left = ls;
const sql_parser::AstNode* probe_key =
probe_key_in_join(join_node->join.condition, probe, build);
if (!probe_key) return nullptr;
const sql_parser::AstNode* build_col =
other_eq_side(join_node->join.condition, probe_key);
if (!build_col) return nullptr;

ScanContext bctx = extract_scan_context(
probe_is_left ? join_node->right : join_node->left);
std::vector<Value> keys = collect_build_join_keys(build, bctx.where_expr, build_col);
if (keys.empty()) return nullptr;

sql_parser::AstNode* in_list = make_in_list_on_column(probe_key, keys);
if (!in_list) return nullptr;

ScanContext pctx = extract_scan_context(
probe_is_left ? join_node->left : join_node->right);
if (!pctx.scan) return nullptr;
const sql_parser::AstNode* probe_where = and_preds(pctx.where_expr, in_list);
PlanNode* probe_dist = distribute_scan(pctx.scan, probe_where,
nullptr, nullptr, nullptr, false);

PlanNode* build_dist = nullptr;
if (bctx.scan && !shards_.is_sharded(build->table_name)) {
sql_parser::StringRef sql = qb_.build_select(
build, bctx.where_expr, nullptr, 0, nullptr, 0,
nullptr, nullptr, 0, -1, false);
build_dist = make_remote_scan(
shards_.get_backend(build->table_name), sql, build);
} else {
build_dist = distribute_node(probe_is_left ? join_node->right : join_node->left);
}
if (!probe_dist || !build_dist) return nullptr;

PlanNode* result = make_plan_node(arena_, PlanNodeType::JOIN);
result->join = join_node->join;
result->left = probe_is_left ? probe_dist : build_dist;
result->right = probe_is_left ? build_dist : probe_dist;
return result;
}

PlanNode* distribute_join(PlanNode* join_node) {
const TableInfo* left_table = find_table(join_node->left);
const TableInfo* right_table = find_table(join_node->right);
Expand All @@ -1217,6 +1374,9 @@ class DistributedPlanner {
return distribute_colocated_join(join_node, left_table, right_table);
}

if (PlanNode* sj = try_semijoin_prune(join_node, left_table, right_table))
return sj;

PlanNode* left_dist = nullptr;
PlanNode* right_dist = nullptr;

Expand Down
2 changes: 1 addition & 1 deletion include/sql_engine/shard_map.h
Original file line number Diff line number Diff line change
Expand Up @@ -177,7 +177,7 @@ class ShardMap {
size_t& out) const {
const TableShardConfig* cfg = lookup(table_name);
if (!cfg || cfg->shards.empty() || !parts || n == 0) return false;
if (n == 1) {
if (n == 1 || cfg->strategy == RoutingStrategy::RANGE) {
return parts[0].is_int
? try_shard_index_for_int(table_name, parts[0].int_val, out)
: try_shard_index_for_string(table_name, parts[0].str,
Expand Down
81 changes: 81 additions & 0 deletions scripts/run_pg_sharding_demo.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
#!/bin/bash
# RANGE + LIST + 2PC against the two PostgreSQL shards.
set -e

SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
PROJECT_DIR="$(dirname "$SCRIPT_DIR")"
cd "$PROJECT_DIR"

if ! docker exec parsersql-pg-shard1 pg_isready -Upostgres &>/dev/null 2>&1; then
echo "ERROR: PG shards not running. Start them with: ./scripts/start_pg_sharding_demo.sh"
exit 1
fi

if [ ! -f ./sqlengine ]; then
echo "Building sqlengine..."
make build-sqlengine
fi

PG1='pgsql://postgres:test@127.0.0.1:16432/testdb?name=pg1'
PG2='pgsql://postgres:test@127.0.0.1:16433/testdb?name=pg2'
TXN_LOG="${TMPDIR:-/tmp}/parsersql-pg-demo.txn"

run_sql() {
local desc="$1"
local sql="$2"
echo "----------------------------------------------"
echo "QUERY: $desc"
echo "SQL: $sql"
echo ""
echo "$sql" | ./sqlengine \
--backend "$PG1" \
--backend "$PG2" \
--shard "users:id:range:5=pg1,10=pg2" \
--shard "regions:name:list:us-east=pg1,us-west=pg2" \
--shard "orders:id:range:105=pg1,110=pg2" \
--txn-log "$TXN_LOG" \
2>&1
echo ""
}

echo "=============================================="
echo " PostgreSQL LIST + RANGE + 2PC demo"
echo "=============================================="
echo " pg1 :16432 users 1-5 / us-east"
echo " pg2 :16433 users 6-10 / us-west"
echo ""

run_sql "RANGE point lookup" \
"SELECT name FROM users WHERE id = 3"

run_sql "RANGE BETWEEN prune" \
"SELECT name FROM users WHERE id BETWEEN 6 AND 10"

run_sql "LIST point lookup" \
"SELECT tz FROM regions WHERE name = 'us-west'"

run_sql "Scatter scan" \
"SELECT COUNT(*) FROM users"

echo "=============================================="
echo " 2PC write across both shards (one engine)"
echo "=============================================="
{
echo "BEGIN"
echo "INSERT INTO users (id, name, age) VALUES (0, 'Zero', 1)"
echo "INSERT INTO users (id, name, age) VALUES (11, 'Eleven', 2)"
echo "COMMIT"
} | ./sqlengine \
--backend "$PG1" \
--backend "$PG2" \
--shard "users:id:range:5=pg1,10=pg2" \
--shard "regions:name:list:us-east=pg1,us-west=pg2" \
--shard "orders:id:range:105=pg1,110=pg2" \
--txn-log "$TXN_LOG" \
2>&1
echo ""

run_sql "Read back 2PC inserts" \
"SELECT id, name FROM users WHERE id IN (0, 11)"

echo "Demo complete. Stop: docker rm -f parsersql-pg-shard1 parsersql-pg-shard2"
98 changes: 98 additions & 0 deletions scripts/start_pg_sharding_demo.sh
Original file line number Diff line number Diff line change
@@ -0,0 +1,98 @@
#!/bin/bash
# Two PostgreSQL shards for LIST + RANGE + 2PC. Ports 16432/16433
# (15432 is the unit-test backend; 13306 is the MySQL sharding demo).
set -e

echo "=== Starting 2-shard PostgreSQL demo ==="

docker rm -f parsersql-pg-shard1 parsersql-pg-shard2 2>/dev/null || true

docker run -d --name parsersql-pg-shard1 \
-p 16432:5432 \
-e POSTGRES_PASSWORD=test \
-e POSTGRES_DB=testdb \
postgres:16 \
-c max_prepared_transactions=16

docker run -d --name parsersql-pg-shard2 \
-p 16433:5432 \
-e POSTGRES_PASSWORD=test \
-e POSTGRES_DB=testdb \
postgres:16 \
-c max_prepared_transactions=16

echo "Waiting for PG shard 1..."
until docker exec parsersql-pg-shard1 pg_isready -Upostgres &>/dev/null 2>&1; do sleep 1; done
echo "PG shard 1 ready"

echo "Waiting for PG shard 2..."
until docker exec parsersql-pg-shard2 pg_isready -Upostgres &>/dev/null 2>&1; do sleep 1; done
echo "PG shard 2 ready"

echo "Loading RANGE users 1-5 + LIST region us-east on shard 1..."
docker exec -i parsersql-pg-shard1 psql -Upostgres testdb <<'SQL'
DROP TABLE IF EXISTS orders;
DROP TABLE IF EXISTS users;
DROP TABLE IF EXISTS regions;

CREATE TABLE users (
id INT PRIMARY KEY,
name VARCHAR(255) NOT NULL,
age INT
);
CREATE TABLE regions (
name VARCHAR(64) PRIMARY KEY,
tz VARCHAR(32)
);
CREATE TABLE orders (
id INT PRIMARY KEY,
user_id INT,
total NUMERIC(10,2)
);

INSERT INTO users VALUES
(1, 'Alice', 30),
(2, 'Bob', 25),
(3, 'Carol', 35),
(4, 'Dave', 28),
(5, 'Eve', 32);
INSERT INTO regions VALUES ('us-east', 'EST');
INSERT INTO orders VALUES (101, 1, 150.00), (102, 3, 50.00);
SQL

echo "Loading RANGE users 6-10 + LIST region us-west on shard 2..."
docker exec -i parsersql-pg-shard2 psql -Upostgres testdb <<'SQL'
DROP TABLE IF EXISTS orders;
DROP TABLE IF EXISTS users;
DROP TABLE IF EXISTS regions;

CREATE TABLE users (
id INT PRIMARY KEY,
name VARCHAR(255) NOT NULL,
age INT
);
CREATE TABLE regions (
name VARCHAR(64) PRIMARY KEY,
tz VARCHAR(32)
);
CREATE TABLE orders (
id INT PRIMARY KEY,
user_id INT,
total NUMERIC(10,2)
);

INSERT INTO users VALUES
(6, 'Frank', 40),
(7, 'Grace', 22),
(8, 'Hank', 31),
(9, 'Ivy', 27),
(10, 'Jack', 36);
INSERT INTO regions VALUES ('us-west', 'PST');
INSERT INTO orders VALUES (106, 6, 80.00), (107, 8, 120.00);
SQL

echo "PostgreSQL shards ready:"
echo " pg1 127.0.0.1:16432 users 1-5, region us-east"
echo " pg2 127.0.0.1:16433 users 6-10, region us-west"
echo "Run: ./scripts/run_pg_sharding_demo.sh"
echo "Stop: docker rm -f parsersql-pg-shard1 parsersql-pg-shard2"
4 changes: 0 additions & 4 deletions src/sql_engine/tool_config_parser.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -208,10 +208,6 @@ ParsedShard parse_shard_spec(const std::string& spec) {
return ps;
}
} else if (strategy_token == "range") {
if (ps.config.shard_key.find('+') != std::string::npos) {
ps.error = "composite shard keys require HASH strategy: " + spec;
return ps;
}
ps.config.strategy = RoutingStrategy::RANGE;
for (auto& entry : split_csv(body)) {
std::string upper_str, backend;
Expand Down
Loading