Skip to content
Open
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
86 changes: 78 additions & 8 deletions dbms/src/Columns/ColumnFunction.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
#include <Columns/ColumnFunction.h>
#include <Columns/countBytesInFilter.h>
#include <Columns/filterColumn.h>
#include <Common/typeid_cast.h>
#include <Functions/IFunction.h>
#include <Interpreters/ExpressionActions.h>
#include <fmt/format.h>
Expand All @@ -28,9 +29,14 @@ namespace ErrorCodes
extern const int LOGICAL_ERROR;
}

ColumnFunction::ColumnFunction(size_t size, FunctionBasePtr function, const ColumnsWithTypeAndName & columns_to_capture)
ColumnFunction::ColumnFunction(
size_t size,
FunctionBasePtr function,
const ColumnsWithTypeAndName & columns_to_capture,
bool is_short_circuit_argument_)
: column_size(size)
, function(function)
, is_short_circuit_argument(is_short_circuit_argument_)
{
appendArguments(columns_to_capture);
}
Expand All @@ -41,7 +47,7 @@ MutableColumnPtr ColumnFunction::cloneResized(size_t size) const
for (auto & column : capture)
column.column = column.column->cloneResized(size);

return ColumnFunction::create(size, function, capture);
return ColumnFunction::create(size, function, capture, is_short_circuit_argument);
}

ColumnPtr ColumnFunction::replicateRange(size_t start_row, size_t end_row, const IColumn::Offsets & offsets) const
Expand All @@ -59,7 +65,7 @@ ColumnPtr ColumnFunction::replicateRange(size_t start_row, size_t end_row, const
column.column = column.column->replicateRange(start_row, end_row, offsets);

size_t replicated_size = 0 == column_size ? 0 : (offsets[end_row - 1]);
return ColumnFunction::create(replicated_size, function, capture);
return ColumnFunction::create(replicated_size, function, capture, is_short_circuit_argument);
}

ColumnPtr ColumnFunction::cut(size_t start, size_t length) const
Expand All @@ -68,7 +74,7 @@ ColumnPtr ColumnFunction::cut(size_t start, size_t length) const
for (auto & column : capture)
column.column = column.column->cut(start, length);

return ColumnFunction::create(length, function, capture);
return ColumnFunction::create(length, function, capture, is_short_circuit_argument);
}

ColumnPtr ColumnFunction::filter(const Filter & filter, ssize_t result_size_hint) const
Expand All @@ -88,7 +94,7 @@ ColumnPtr ColumnFunction::filter(const Filter & filter, ssize_t result_size_hint
else
filtered_size = capture.front().column->size();

return ColumnFunction::create(filtered_size, function, capture);
return ColumnFunction::create(filtered_size, function, capture, is_short_circuit_argument);
}

ColumnPtr ColumnFunction::permute(const Permutation & perm, size_t limit) const
Expand All @@ -107,7 +113,7 @@ ColumnPtr ColumnFunction::permute(const Permutation & perm, size_t limit) const
for (auto & column : capture)
column.column = column.column->permute(perm, limit);

return ColumnFunction::create(limit, function, capture);
return ColumnFunction::create(limit, function, capture, is_short_circuit_argument);
}

std::vector<MutableColumnPtr> ColumnFunction::scatter(
Expand Down Expand Up @@ -138,7 +144,7 @@ std::vector<MutableColumnPtr> ColumnFunction::scatter(
{
auto & capture = captures[part];
size_t s = capture.empty() ? counts[part] : capture.front().column->size();
columns.emplace_back(ColumnFunction::create(s, function, std::move(capture)));
columns.emplace_back(ColumnFunction::create(s, function, std::move(capture), is_short_circuit_argument));
}

return columns;
Expand Down Expand Up @@ -253,7 +259,19 @@ ColumnWithTypeAndName ColumnFunction::reduce() const
captured),
ErrorCodes::LOGICAL_ERROR);

Block block(captured_columns);
if (is_short_circuit_argument && column_size == 0)
return {function->getReturnType()->createColumn(), function->getReturnType(), ""};

auto columns = captured_columns;
if (is_short_circuit_argument)
{
const size_t required_arguments = function->isShortCircuit() ? 1 : columns.size();
for (size_t i = 0; i < required_arguments; ++i)
if (const auto * deferred = checkAndGetShortCircuitArgument(columns[i].column))
columns[i].column = deferred->reduce().column;
}

Block block(columns);
block.insert({nullptr, function->getReturnType(), ""});

ColumnNumbers arguments(captured_columns.size());
Expand All @@ -265,4 +283,56 @@ ColumnWithTypeAndName ColumnFunction::reduce() const
return block.getByPosition(captured_columns.size());
}

const ColumnFunction * checkAndGetShortCircuitArgument(const ColumnPtr & column)
{
const auto * function = typeid_cast<const ColumnFunction *>(column.get());
return function && function->isShortCircuitArgument() ? function : nullptr;
}

void maskedExecute(ColumnWithTypeAndName & column, const IColumn::Filter & mask)
{
const auto * deferred = checkAndGetShortCircuitArgument(column.column);
if (!deferred)
return;

RUNTIME_CHECK(column.column->size() == mask.size());
const size_t selected = countBytesInFilter(mask);
if (selected == 0)
{
column.column = column.type->createColumnConstWithDefaultValue(mask.size());
return;
}
if (selected == mask.size())
{
column.column = deferred->reduce().column;
return;
}

auto filtered = deferred->filter(mask, selected);
auto result = static_cast<const ColumnFunction &>(*filtered).reduce().column;
if (auto materialized = result->convertToFullColumnIfConst())
result = std::move(materialized);
RUNTIME_CHECK(result->size() == selected);

// Use the existing bulk insertion interface instead of adding expand() to every TiFlash column type.
auto expanded = result->cloneEmpty();
expanded->reserve(mask.size());
size_t source = 0;
for (size_t begin = 0; begin < mask.size();)
{
size_t end = begin + 1;
while (end < mask.size() && (mask[end] != 0) == (mask[begin] != 0))
++end;
if (mask[begin])
{
expanded->insertRangeFrom(*result, source, end - begin);
source += end - begin;
}
else
expanded->insertManyDefaults(end - begin);
begin = end;
}
column.column = std::move(expanded);
}

} // namespace DB
17 changes: 14 additions & 3 deletions dbms/src/Columns/ColumnFunction.h
Original file line number Diff line number Diff line change
Expand Up @@ -25,15 +25,19 @@ namespace DB
class IFunctionBase;
using FunctionBasePtr = std::shared_ptr<IFunctionBase>;

/** A column containing a lambda expression.
* Behaves like a constant-column. Contains an expression, but not input or output data.
/** A column containing a lambda or deferred scalar expression and its captured arguments.
* A deferred scalar expression is evaluated only after its captures have been filtered.
*/
class ColumnFunction final : public COWPtrHelper<IColumn, ColumnFunction>
{
private:
friend class COWPtrHelper<IColumn, ColumnFunction>;

ColumnFunction(size_t size, FunctionBasePtr function, const ColumnsWithTypeAndName & columns_to_capture);
ColumnFunction(
size_t size,
FunctionBasePtr function,
const ColumnsWithTypeAndName & columns_to_capture,
bool is_short_circuit_argument = false);

public:
const char * getFamilyName() const override { return "Function"; }
Expand Down Expand Up @@ -70,6 +74,7 @@ class ColumnFunction final : public COWPtrHelper<IColumn, ColumnFunction>

void appendArguments(const ColumnsWithTypeAndName & columns);
ColumnWithTypeAndName reduce() const;
bool isShortCircuitArgument() const { return is_short_circuit_argument; }

Field operator[](size_t) const override
{
Expand Down Expand Up @@ -283,8 +288,14 @@ class ColumnFunction final : public COWPtrHelper<IColumn, ColumnFunction>
size_t column_size;
FunctionBasePtr function;
ColumnsWithTypeAndName captured_columns;
bool is_short_circuit_argument;

void appendArgument(const ColumnWithTypeAndName & column);
};

const ColumnFunction * checkAndGetShortCircuitArgument(const ColumnPtr & column);

/// Adapted from ClickHouse Columns/MaskOperations.cpp: filter captures, reduce, then restore row positions.
void maskedExecute(ColumnWithTypeAndName & column, const IColumn::Filter & mask);

} // namespace DB
38 changes: 0 additions & 38 deletions dbms/src/Flash/Coprocessor/DAGExpressionAnalyzer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -49,8 +49,6 @@
#include <tipb/executor.pb.h>
#include <tipb/expression.pb.h>

#include <ext/scope_guard.h>

namespace DB
{
namespace ErrorCodes
Expand Down Expand Up @@ -1003,13 +1001,6 @@ String DAGExpressionAnalyzer::buildFilterColumn(
const google::protobuf::RepeatedPtrField<tipb::Expr> & conditions,
bool null_as_false)
{
building_filter_conditions = true;
json_valid_guarded_exprs.clear();
SCOPE_EXIT({
building_filter_conditions = false;
json_valid_guarded_exprs.clear();
});

String filter_column_name;
if (conditions.size() == 1)
{
Expand All @@ -1030,12 +1021,7 @@ String DAGExpressionAnalyzer::buildFilterColumn(
{
Names arg_names;
for (const auto & condition : conditions)
{
auto guards_before_condition = json_valid_guarded_exprs;
arg_names.push_back(getActions(condition, actions, true));
json_valid_guarded_exprs = std::move(guards_before_condition);
recordJsonValidGuards(condition);
}
// connect all the conditions by logical and
// two_value_and treats null as false inside the `two_value_and` function, so the output column
// will always be UInt8 type, which can save the merge step in FilterDescription
Expand All @@ -1046,30 +1032,6 @@ String DAGExpressionAnalyzer::buildFilterColumn(
return filter_column_name;
}

void DAGExpressionAnalyzer::recordJsonValidGuards(const tipb::Expr & expr)
{
if (!building_filter_conditions || !isScalarFunctionExpr(expr))
return;

if (expr.sig() == tipb::ScalarFuncSig::JsonValidStringSig && expr.children_size() == 1)
{
json_valid_guarded_exprs.emplace(exprToString(expr.children(0), getCurrentInputColumns()));
return;
}

if (expr.sig() == tipb::ScalarFuncSig::LogicalAnd)
{
for (const auto & child : expr.children())
recordJsonValidGuards(child);
}
}

bool DAGExpressionAnalyzer::isJsonValidGuarded(const tipb::Expr & expr) const
{
return building_filter_conditions
&& json_valid_guarded_exprs.contains(exprToString(expr, getCurrentInputColumns()));
}

std::tuple<ExpressionActionsPtr, String, ExpressionActionsPtr> DAGExpressionAnalyzer::buildPushDownFilter(
const google::protobuf::RepeatedPtrField<tipb::Expr> & conditions,
bool null_as_false)
Expand Down
6 changes: 0 additions & 6 deletions dbms/src/Flash/Coprocessor/DAGExpressionAnalyzer.h
Original file line number Diff line number Diff line change
Expand Up @@ -320,17 +320,11 @@ class DAGExpressionAnalyzer : private boost::noncopyable
const std::vector<tipb::FieldType> & require_schema,
const std::vector<Int32> & output_offsets) const;

void recordJsonValidGuards(const tipb::Expr & expr);
bool isJsonValidGuarded(const tipb::Expr & expr) const;

// all columns from table scan
NamesAndTypes source_columns;
DAGPreparedSets prepared_sets;
const Context & context;

bool building_filter_conditions = false;
std::unordered_set<String> json_valid_guarded_exprs;

friend class DAGExpressionAnalyzerHelper;
};

Expand Down
15 changes: 1 addition & 14 deletions dbms/src/Flash/Coprocessor/DAGExpressionAnalyzerHelper.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -202,20 +202,13 @@ String DAGExpressionAnalyzerHelper::buildLogicalFunction(
const ExpressionActionsPtr & actions)
{
const String & func_name = getFunctionName(expr);
auto guards_before_function = analyzer->json_valid_guarded_exprs;
Names argument_names;
for (const auto & child : expr.children())
{
auto guards_before_child = analyzer->json_valid_guarded_exprs;
String name = analyzer->getActions(child, actions, true);
argument_names.push_back(name);
analyzer->json_valid_guarded_exprs = std::move(guards_before_child);
if (func_name == "and" || func_name == "two_value_and")
analyzer->recordJsonValidGuards(child);
}
String result = analyzer->applyFunction(func_name, argument_names, actions, getCollatorFromExpr(expr));
analyzer->json_valid_guarded_exprs = std::move(guards_before_function);
return result;
return analyzer->applyFunction(func_name, argument_names, actions, getCollatorFromExpr(expr));
}

// left(str,len) = substrUTF8(str,1,len)
Expand Down Expand Up @@ -306,12 +299,7 @@ String DAGExpressionAnalyzerHelper::buildSingleParamJsonRelatedFunctions(
const auto & input_expr = expr.children(0);
String arg = analyzer->getActions(input_expr, actions);
const auto & collator = getCollatorFromExpr(expr);
const bool ignore_invalid_json
= func_name == FunctionCastStringAsJson::name && analyzer->isJsonValidGuarded(input_expr);
String result_name = genFuncString(func_name, {arg}, {collator}, {&input_expr.field_type(), &expr.field_type()});
// Guarded and strict casts can coexist in different logical branches and must not share an action.
if (ignore_invalid_json)
result_name += "_json_valid_guarded";
if (actions->getSampleBlock().has(result_name))
return result_name;

Expand All @@ -330,7 +318,6 @@ String DAGExpressionAnalyzerHelper::buildSingleParamJsonRelatedFunctions(
{
function_cast_string_as_json->setInputTiDBFieldType(input_expr.field_type());
function_cast_string_as_json->setOutputTiDBFieldType(expr.field_type());
function_cast_string_as_json->setIgnoreInvalidJson(ignore_invalid_json);
}
else if (auto * function_cast_time_as_json = dynamic_cast<FunctionCastTimeAsJson *>(function_impl);
function_cast_time_as_json)
Expand Down
37 changes: 37 additions & 0 deletions dbms/src/Flash/tests/gtest_filter_executor.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,43 @@ try
}
CATCH

TEST_F(FilterExecutorTestRunner, ShortCircuitJsonGuard)
try
{
context.addMockTable(
{"test_db", "json_guard"},
{{"document", TiDB::TP::TypeString}},
{toNullableVec<String>("document", {"", "invalid json", R"({"a": 1})", {}, R"({"b": 2})"})});
auto request = context.scan("test_db", "json_guard").filter(eq(col("document"), col("document"))).build(context);
auto * executor = request->has_root_executor() ? request->mutable_root_executor()
: request->mutable_executors(request->executors_size() - 1);
ASSERT_TRUE(executor->has_selection());
auto * selection = executor->mutable_selection();
const auto column_ref = selection->conditions(0).children(0);
auto json_valid = selection->conditions(0);
json_valid.set_sig(tipb::ScalarFuncSig::JsonValidStringSig);
json_valid.clear_children();
*json_valid.add_children() = column_ref;
auto cast_json = json_valid;
cast_json.set_sig(tipb::ScalarFuncSig::CastStringAsJson);
cast_json.mutable_field_type()->set_tp(TiDB::TypeJSON);
cast_json.mutable_field_type()->set_flag(TiDB::ColumnFlagParseToJSON);
auto is_null = json_valid;
is_null.set_sig(tipb::ScalarFuncSig::StringIsNull);
*is_null.mutable_children(0) = cast_json;
auto is_not_null = json_valid;
is_not_null.set_sig(tipb::ScalarFuncSig::UnaryNotInt);
*is_not_null.mutable_children(0) = is_null;
selection->clear_conditions();
*selection->add_conditions() = json_valid;
*selection->add_conditions() = is_not_null;

WRAP_FOR_TEST_BEGIN
executeAndAssertColumnsEqual(request, {toNullableVec<String>({R"({"a": 1})", R"({"b": 2})"})});
WRAP_FOR_TEST_END
}
CATCH

TEST_F(FilterExecutorTestRunner, andOr)
try
{
Expand Down
2 changes: 2 additions & 0 deletions dbms/src/Functions/FunctionsConversion.h
Original file line number Diff line number Diff line change
Expand Up @@ -2650,6 +2650,8 @@ class ExecutableFunctionCast : public IExecutableFunction
class FunctionCast final : public IFunctionBase
{
public:
bool isSuitableForShortCircuitArgumentsExecution() const override { return true; }

using WrapperType = std::function<void(Block &, const ColumnNumbers &, size_t)>;
using MonotonicityForRange = std::function<Monotonicity(const IDataType &, const Field &, const Field &)>;

Expand Down
Loading