summaryrefslogtreecommitdiffstats
path: root/yql/essentials/minikql/comp_nodes/mkql_block_variant_item.cpp
blob: 82db82fa17b9ca23bb976de22da2bba9b18370e4 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
#include "mkql_block_variant_item.h"

#include <yql/essentials/minikql/computation/mkql_block_impl.h>
#include <yql/essentials/minikql/computation/mkql_block_reader.h>
#include <yql/essentials/minikql/computation/mkql_computation_node_holders.h>
#include <yql/essentials/minikql/mkql_node_builder.h>
#include <yql/essentials/minikql/mkql_node_cast.h>
#include <yql/essentials/public/udf/arrow/block_builder.h>
#include <yql/essentials/public/udf/arrow/block_reader.h>

namespace NKikimr::NMiniKQL {

namespace {

template <bool IsOptional>
class TVariantItemBlockExec {
public:
    class TVariantItemKernelState: public arrow::compute::KernelState {
    public:
        explicit TVariantItemKernelState(TType* inputItemType)
            : Reader_(MakeBlockReader(TTypeInfoHelper(), inputItemType))
        {
        }

        IBlockReader& GetReader() {
            return *Reader_;
        }

    private:
        std::unique_ptr<IBlockReader> Reader_;
    };

    explicit TVariantItemBlockExec(TType* inputItemType, TType* resultItemType)
        : InputItemType_(inputItemType)
        , ResultItemType_(resultItemType)
    {
    }

    arrow::Status Exec(arrow::compute::KernelContext* ctx, const arrow::compute::ExecBatch& batch, arrow::Datum* res) const {
        auto& reader = static_cast<TVariantItemKernelState&>(*ctx->state()).GetReader();
        const arrow::Datum& variantDatum = batch.values[0];

        if (variantDatum.is_scalar()) {
            *res = ConvertScalar(ResultItemType_,
                                 ComputeOutputItem(reader.GetScalarItem(*variantDatum.scalar())),
                                 *ctx->memory_pool());
            return arrow::Status::OK();
        }

        MKQL_ENSURE(variantDatum.is_array(), "Expected array datum");
        const auto& variantArrayData = variantDatum.array();
        const size_t length = static_cast<size_t>(variantArrayData->length);
        auto builder = NYql::NUdf::MakeArrayBuilder(TTypeInfoHelper(), ResultItemType_, *ctx->memory_pool(), length, /*pgBuilder=*/nullptr);
        for (size_t i = 0; i < length; ++i) {
            builder->Add(ComputeOutputItem(reader.GetItem(*variantArrayData, i)));
        }
        *res = builder->Build(/*finish=*/true);
        return arrow::Status::OK();
    }

private:
    TBlockItem ComputeOutputItem(TBlockItem blockItem) const {
        if constexpr (IsOptional) {
            if (!blockItem) {
                return TBlockItem{};
            }
            return blockItem.GetVariantItem().MakeOptional();
        } else {
            return blockItem.GetVariantItem();
        }
    }

    TType* const InputItemType_;
    TType* const ResultItemType_;
};

template <bool IsOptional>
std::shared_ptr<arrow::compute::ScalarKernel> MakeBlockVariantItemKernel(const TVector<TType*>& argTypes,
                                                                         TType* resultType,
                                                                         TType* inputItemType) {
    using TExec = TVariantItemBlockExec<IsOptional>;
    auto exec = std::make_shared<TExec>(
        inputItemType,
        AS_TYPE(TBlockType, resultType)->GetItemType());
    auto kernel = std::make_shared<arrow::compute::ScalarKernel>(
        ConvertToInputTypes(argTypes),
        ConvertToOutputType(resultType),
        [exec](arrow::compute::KernelContext* ctx, const arrow::compute::ExecBatch& batch, arrow::Datum* res) {
            return exec->Exec(ctx, batch, res);
        });
    kernel->null_handling = arrow::compute::NullHandling::COMPUTED_NO_PREALLOCATE;
    kernel->mem_allocation = arrow::compute::MemAllocation::NO_PREALLOCATE;
    kernel->init = [inputItemType](arrow::compute::KernelContext*, const arrow::compute::KernelInitArgs&) {
        return arrow::Result(std::make_unique<typename TExec::TVariantItemKernelState>(inputItemType));
    };
    return kernel;
}

} // namespace

IComputationNode* WrapBlockVariantItem(TCallable& callable, const TComputationNodeFactoryContext& ctx) {
    MKQL_ENSURE(callable.GetInputsCount() == 1, "Expected 1 argument");

    auto blockType = AS_TYPE(TBlockType, callable.GetInput(0).GetStaticType());
    auto inputItemType = blockType->GetItemType();

    bool isOptional;
    auto variantItemType = UnpackOptional(inputItemType, isOptional);
    AS_TYPE(TVariantType, variantItemType);

    auto variantCompute = LocateNode(ctx.NodeLocator, callable, 0);
    TComputationNodePtrVector argsNodes = {variantCompute};
    TVector<TType*> argsTypes = {blockType};

    auto resultType = callable.GetType()->GetReturnType();

    auto kernel = isOptional
                      ? MakeBlockVariantItemKernel<true>(argsTypes, resultType, inputItemType)
                      : MakeBlockVariantItemKernel<false>(argsTypes, resultType, inputItemType);

    return new TBlockFuncNode(ctx.Mutables, ctx.RuntimeSettings->DatumValidation.Get(),
                              callable.GetType()->GetName(), std::move(argsNodes), argsTypes, resultType, *kernel, kernel);
}

} // namespace NKikimr::NMiniKQL