#include "lbug_partition_routing.h"
#include <algorithm>
#include <cstring>
#include <map>
#include <memory>
#include <mutex>
#include <shared_mutex>
#include <stdexcept>
#include <string>
#include <utility>
#include <vector>
#include <lbug.hpp>
namespace {
using lbug::common::LogicalType;
using lbug::common::LogicalTypeID;
using lbug::common::Value;
using lbug::common::ValueVector;
using lbug::function::TableFuncBindData;
using lbug::function::TableFuncBindInput;
using lbug::function::TableFuncInput;
using lbug::function::TableFuncMorsel;
using lbug::function::TableFunction;
struct ParentSchema {
std::vector<std::string> propNames;
std::vector<LogicalTypeID> propTypeIDs;
};
struct Registry {
std::shared_mutex mutex;
lbug_partition_hooks_t callbacks{};
std::map<uint64_t, std::vector<std::vector<Value>>> rows;
std::map<uint64_t, ParentSchema> schemas;
std::map<uint64_t, std::unique_ptr<TableFunction>> scanFunctions;
};
Registry& registry() {
static Registry instance;
return instance;
}
lbug_partition_ref_t toCRef(lbug::common::PartitionRef ref) {
return lbug_partition_ref_t{ref.parentTableID, ref.partitionIndex};
}
std::vector<Value> materializeRow(std::span<ValueVector* const> vectors, uint32_t selPos) {
std::vector<Value> cells;
cells.reserve(vectors.size());
for (const ValueVector* vec : vectors) {
if (vec == nullptr || vec->isNull(selPos)) {
cells.push_back(Value::createNullValue());
} else {
cells.push_back(*vec->getAsValue(selPos));
}
}
return cells;
}
std::vector<const void*> borrowCells(const std::vector<Value>& cells) {
std::vector<const void*> ptrs;
ptrs.reserve(cells.size());
for (const Value& cell : cells) {
ptrs.push_back(static_cast<const void*>(&cell));
}
return ptrs;
}
bool locateForward(void* context, lbug::common::PartitionRef ref, void** handleOut) {
const auto& cbs = registry().callbacks;
if (cbs.locate == nullptr) {
return false;
}
return cbs.locate(context, toCRef(ref), handleOut) != 0;
}
void createForward(void* context, lbug::common::PartitionRef ref, void* handle) {
const auto& cbs = registry().callbacks;
if (cbs.on_partition_create != nullptr) {
cbs.on_partition_create(context, toCRef(ref), handle);
}
}
void dropForward(void* context, lbug::common::PartitionRef ref, void* handle) {
const auto& cbs = registry().callbacks;
if (cbs.on_partition_drop != nullptr) {
cbs.on_partition_drop(context, toCRef(ref), handle);
}
}
lbug::common::nodeID_t storeAndForward(lbug::common::PartitionRef ref, void* handle,
std::vector<Value> cells) {
auto& reg = registry();
uint64_t offset = 0;
std::unique_ptr<std::vector<Value>> snapshot;
{
std::unique_lock lock(reg.mutex);
auto& store = reg.rows[ref.parentTableID];
offset = store.size();
store.push_back(std::move(cells));
snapshot = std::make_unique<std::vector<Value>>(store.back());
}
const auto ptrs = borrowCells(*snapshot);
const auto& cbs = reg.callbacks;
if (cbs.insert_row != nullptr) {
cbs.insert_row(reg.callbacks.context, toCRef(ref), handle, ptrs.data(), ptrs.size());
}
return lbug::common::nodeID_t{static_cast<lbug::common::offset_t>(offset),
ref.parentTableID};
}
lbug::common::nodeID_t insertRowForward(void* context, lbug::common::PartitionRef ref,
void* handle, lbug::transaction::Transaction* , const ValueVector* keyVector,
std::span<ValueVector* const> columnVectors) {
(void)context;
(void)keyVector;
const uint32_t selPos = columnVectors.empty() ?
0 :
columnVectors[0]->state->getSelVector()[0];
return storeAndForward(ref, handle, materializeRow(columnVectors, selPos));
}
void insertChunkForward(void* context, lbug::common::PartitionRef ref, void* handle,
lbug::transaction::Transaction* , const ValueVector* keyVector,
std::span<ValueVector* const> columnVectors, uint64_t startRow, uint64_t numRows) {
(void)context;
(void)keyVector;
const auto& sel = columnVectors[0]->state->getSelVector();
for (uint64_t j = 0; j < numRows; ++j) {
storeAndForward(ref, handle, materializeRow(columnVectors, sel[startRow + j]));
}
}
lbug::binder::expression_vector scanColumns(
const ParentSchema& schema, const std::string& nodeUniqueName) {
lbug::binder::expression_vector columns;
columns.push_back(std::make_shared<lbug::binder::VariableExpression>(
LogicalType(LogicalTypeID::INT64), nodeUniqueName + "._ID", "rowid"));
for (size_t i = 0; i < schema.propNames.size(); ++i) {
columns.push_back(std::make_shared<lbug::binder::VariableExpression>(
LogicalType(schema.propTypeIDs[i]), nodeUniqueName + "." + schema.propNames[i],
schema.propNames[i]));
}
return columns;
}
lbug::common::row_idx_t storeSize(uint64_t parentTableID) {
auto& reg = registry();
std::shared_lock lock(reg.mutex);
const auto it = reg.rows.find(parentTableID);
return it == reg.rows.end() ?
0 :
static_cast<lbug::common::row_idx_t>(it->second.size());
}
lbug::common::offset_t scanParent(
uint64_t parentTableID, const TableFuncMorsel& morsel, lbug::common::DataChunk& output) {
if (!morsel.hasMoreToOutput()) {
return 0;
}
auto& reg = registry();
std::vector<std::vector<Value>> slice;
uint64_t start = 0;
uint64_t count = 0;
{
std::shared_lock lock(reg.mutex);
const auto it = reg.rows.find(parentTableID);
const uint64_t storeSize =
(it == reg.rows.end()) ? 0 : static_cast<uint64_t>(it->second.size());
start = static_cast<uint64_t>(morsel.startOffset);
if (start >= storeSize) {
return 0;
}
count = std::min(morsel.getMorselSize(), storeSize - start);
std::vector<std::vector<Value>> fresh(
it->second.begin() + start, it->second.begin() + start + count);
slice.swap(fresh);
}
const uint64_t numOutputCols = output.getNumValueVectors();
for (uint64_t i = 0; i < count; ++i) {
output.getValueVectorMutable(0).copyFromValue(
i, Value(static_cast<int64_t>(start + i)));
const auto& row = slice[i];
for (uint64_t c = 1; c < numOutputCols; ++c) {
auto& vec = output.getValueVectorMutable(c);
if (c - 1 >= row.size() || row[c - 1].isNull()) {
vec.setNull(static_cast<uint32_t>(i), true);
} else {
vec.copyFromValue(i, row[c - 1]);
}
}
}
return static_cast<lbug::common::offset_t>(count);
}
bool bindScanForward(void* context, lbug::common::PartitionRef ref, void* ,
lbug::common::PartitionScanSpec* specOut) {
(void)context;
if (specOut == nullptr) {
return false;
}
auto& reg = registry();
lbug::function::TableFunction* scanFunction = nullptr;
std::shared_ptr<ParentSchema> schema;
{
std::unique_lock lock(reg.mutex);
const auto schemaIt = reg.schemas.find(ref.parentTableID);
if (schemaIt == reg.schemas.end()) {
return false;
}
schema = std::make_shared<ParentSchema>(schemaIt->second);
auto funcIt = reg.scanFunctions.find(ref.parentTableID);
if (funcIt == reg.scanFunctions.end()) {
auto func = std::make_unique<lbug::function::TableFunction>(
"lbug_routing_scan_" + std::to_string(ref.parentTableID),
std::vector<LogicalTypeID>{});
func->bindFunc = [schema](lbug::main::ClientContext*,
const TableFuncBindInput*) {
return std::make_unique<TableFuncBindData>(scanColumns(*schema, "r"), 0);
};
const uint64_t parent = ref.parentTableID;
func->tableFunc = lbug::function::SimpleTableFunc::getTableFunc(
[parent](const TableFuncMorsel& morsel, const TableFuncInput&,
lbug::common::DataChunk& output) {
return scanParent(parent, morsel, output);
});
func->initSharedStateFunc = lbug::function::SimpleTableFunc::initSharedState;
func->initLocalStateFunc = lbug::function::TableFunction::initEmptyLocalState;
scanFunction = func.get();
reg.scanFunctions.emplace(ref.parentTableID, std::move(func));
} else {
scanFunction = funcIt->second.get();
}
}
specOut->scanFunction = scanFunction;
const uint64_t parent = ref.parentTableID;
specOut->createBindData = [schema, parent](const std::string& nodeUniqueName) {
return std::make_unique<TableFuncBindData>(
scanColumns(*schema, nodeUniqueName), storeSize(parent));
};
return true;
}
}
extern "C" {
int lbug_partition_routing_install(const lbug_partition_hooks_t* hooks) {
if (hooks == nullptr) {
return 1;
}
if (lbug::common::getPartitionRoutingHooks() != nullptr) {
return 1;
}
auto& reg = registry();
{
std::unique_lock lock(reg.mutex);
reg.callbacks = *hooks;
}
static lbug::common::PartitionRoutingHooks engineHooks{};
engineHooks.context = hooks->context;
engineHooks.locate = locateForward;
engineHooks.onPartitionCreate = createForward;
engineHooks.onPartitionDrop = dropForward;
engineHooks.bindScan = bindScanForward;
engineHooks.insertRow = insertRowForward;
engineHooks.insertChunk = insertChunkForward;
engineHooks.lookupRow = nullptr;
lbug::common::setPartitionRoutingHooks(&engineHooks);
return 0;
}
void lbug_partition_routing_uninstall(void) {
lbug::common::setPartitionRoutingHooks(nullptr);
}
uint8_t lbug_partition_routing_is_installed(void) {
return lbug::common::getPartitionRoutingHooks() != nullptr ? 1 : 0;
}
int lbug_partition_routing_register_schema(uint64_t parent_table_id,
const char* const* prop_names, const uint8_t* type_ids, size_t n_props) {
if ((n_props > 0 && (prop_names == nullptr || type_ids == nullptr)) || n_props == 0) {
return 1;
}
ParentSchema schema;
try {
for (size_t i = 0; i < n_props; ++i) {
if (prop_names[i] == nullptr) {
return 1;
}
schema.propNames.emplace_back(prop_names[i]);
schema.propTypeIDs.push_back(static_cast<LogicalTypeID>(type_ids[i]));
}
} catch (...) {
return 1;
}
auto& reg = registry();
std::unique_lock lock(reg.mutex);
reg.schemas[parent_table_id] = std::move(schema);
return 0;
}
}