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
1 change: 1 addition & 0 deletions extensions/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -470,6 +470,7 @@ cc_test(
":lists_functions",
"//checker:type_check_issue",
"//checker:validation_result",
"//common:ast",
"//common:source",
"//common:value",
"//common:value_testing",
Expand Down
73 changes: 53 additions & 20 deletions extensions/lists_functions.cc
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,8 @@ absl::Span<const cel::Type> SortableTypes() {
return kTypes;
}

constexpr int64_t kMaxRangeSize = 1000000;

// Slow distinct() implementation that uses Equal() to compare values in O(n^2).
absl::Status ListDistinctHeterogeneousImpl(
const ListValue& list,
Expand Down Expand Up @@ -223,10 +225,20 @@ absl::StatusOr<Value> ListFlatten(
return std::move(*builder).Build();
}

absl::StatusOr<ListValue> ListRange(
int64_t end, const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
absl::StatusOr<Value> ListRange(
int64_t end, int64_t max_range_size,
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
google::protobuf::MessageFactory* absl_nonnull message_factory,
google::protobuf::Arena* absl_nonnull arena) {
if (end < 0) {
return ErrorValue(absl::InvalidArgumentError(absl::StrFormat(
"lists.range: size must be non-negative, got %d", end)));
}
if (end > max_range_size) {
return ErrorValue(absl::InvalidArgumentError(
absl::StrFormat("lists.range: size %d exceeds maximum allowed (%d)",
end, max_range_size)));
}
auto builder = NewListValueBuilder(arena);
builder->Reserve(end);
for (int64_t i = 0; i < end; ++i) {
Expand Down Expand Up @@ -512,11 +524,27 @@ absl::Status RegisterListFlattenFunction(FunctionRegistry& registry) {
return absl::OkStatus();
}

absl::Status RegisterListRangeFunction(FunctionRegistry& registry) {
return UnaryFunctionAdapter<absl::StatusOr<Value>,
int64_t>::RegisterGlobalOverload("lists.range",
&ListRange,
registry);
absl::Status RegisterListRangeFunction(
FunctionRegistry& registry,
const ListsExtensionOptions& extension_options) {
constexpr int64_t kMaxRangeSize = 1000000;
int64_t effective_limit = kMaxRangeSize;
if (extension_options.max_range_size > 0 &&
extension_options.max_range_size < effective_limit) {
effective_limit = extension_options.max_range_size;
}
return UnaryFunctionAdapter<absl::StatusOr<Value>, int64_t>::
RegisterGlobalOverload(
"lists.range",
[effective_limit](
int64_t end,
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
google::protobuf::MessageFactory* absl_nonnull message_factory,
google::protobuf::Arena* absl_nonnull arena) -> absl::StatusOr<Value> {
return ListRange(end, effective_limit, descriptor_pool,
message_factory, arena);
},
registry);
}

absl::Status RegisterListReverseFunction(FunctionRegistry& registry) {
Expand Down Expand Up @@ -657,23 +685,23 @@ absl::Status ConfigureParser(ParserBuilder& builder, int version) {

} // namespace

absl::Status RegisterListsFunctions(FunctionRegistry& registry,
const RuntimeOptions& options,
int version) {
absl::Status RegisterListsFunctions(
FunctionRegistry& registry, const RuntimeOptions& options,
const ListsExtensionOptions& extension_options) {
CEL_RETURN_IF_ERROR(RegisterListSliceFunction(registry));
if (version == 0) {
if (extension_options.version == 0) {
return absl::OkStatus();
}

// Since version 1
CEL_RETURN_IF_ERROR(RegisterListFlattenFunction(registry));
if (version == 1) {
if (extension_options.version == 1) {
return absl::OkStatus();
}

// Since version 2
CEL_RETURN_IF_ERROR(RegisterListDistinctFunction(registry));
CEL_RETURN_IF_ERROR(RegisterListRangeFunction(registry));
CEL_RETURN_IF_ERROR(RegisterListRangeFunction(registry, extension_options));
CEL_RETURN_IF_ERROR(RegisterListReverseFunction(registry));
CEL_RETURN_IF_ERROR(RegisterListSortFunction(registry));
return absl::OkStatus();
Expand All @@ -684,18 +712,23 @@ absl::Status RegisterListsMacros(MacroRegistry& registry, const ParserOptions&,
return registry.RegisterMacros(lists_macros(version));
}

CheckerLibrary ListsCheckerLibrary(int version) {
CheckerLibrary ListsCheckerLibrary(
const ListsExtensionOptions& extension_options) {
return {.id = "cel.lib.ext.lists",
.configure = [version](TypeCheckerBuilder& builder) {
.configure = [version = extension_options.version](
TypeCheckerBuilder& builder) {
return RegisterListsCheckerDecls(builder, version);
}};
}

CompilerLibrary ListsCompilerLibrary(int version) {
auto lib = CompilerLibrary::FromCheckerLibrary(ListsCheckerLibrary(version));
lib.configure_parser = [version](ParserBuilder& builder) {
return ConfigureParser(builder, version);
};
CompilerLibrary ListsCompilerLibrary(
const ListsExtensionOptions& extension_options) {
auto lib = CompilerLibrary::FromCheckerLibrary(
ListsCheckerLibrary(extension_options));
lib.configure_parser =
[version = extension_options.version](ParserBuilder& builder) {
return ConfigureParser(builder, version);
};
return lib;
}

Expand Down
47 changes: 41 additions & 6 deletions extensions/lists_functions.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,8 @@
#ifndef THIRD_PARTY_CEL_CPP_EXTENSIONS_LISTS_FUNCTIONS_H_
#define THIRD_PARTY_CEL_CPP_EXTENSIONS_LISTS_FUNCTIONS_H_

#include <cstdint>

#include "absl/status/status.h"
#include "checker/type_checker_builder.h"
#include "compiler/compiler.h"
Expand All @@ -27,6 +29,18 @@ namespace cel::extensions {

constexpr int kListsExtensionLatestVersion = 2;

struct ListsExtensionOptions {
int version = kListsExtensionLatestVersion;

// Maximum size allowed for lists.range().
// Setting a tighter limit (e.g. 100) will restrict the max size further.
// A standard limit of 1,000,000 applies if a tighter limit isn't
// configured.
int64_t max_range_size = 1000000;
};

using ListsFunctionsOptions = ListsExtensionOptions;

// Register implementations for list extension functions.
//
// === Since version 0 ===
Expand All @@ -45,9 +59,17 @@ constexpr int kListsExtensionLatestVersion = 2;
//
// <list(T)>.sort() -> list(T)
//
absl::Status RegisterListsFunctions(FunctionRegistry& registry,
const RuntimeOptions& options,
int version = kListsExtensionLatestVersion);
absl::Status RegisterListsFunctions(
FunctionRegistry& registry, const RuntimeOptions& options,
const ListsExtensionOptions& extension_options = {});

inline absl::Status RegisterListsFunctions(FunctionRegistry& registry,
const RuntimeOptions& options,
int version) {
ListsExtensionOptions extension_options;
extension_options.version = version;
return RegisterListsFunctions(registry, options, extension_options);
}

// Register list macros.
//
Expand Down Expand Up @@ -76,7 +98,14 @@ absl::Status RegisterListsMacros(MacroRegistry& registry,
// <list(T)>.reverse() -> list(T)
//
// <list(T_)>.sort() -> list(T_) where T_ is partially orderable
CheckerLibrary ListsCheckerLibrary(int version = kListsExtensionLatestVersion);
CheckerLibrary ListsCheckerLibrary(
const ListsExtensionOptions& extension_options = {});

inline CheckerLibrary ListsCheckerLibrary(int version) {
ListsExtensionOptions extension_options;
extension_options.version = version;
return ListsCheckerLibrary(extension_options);
}

// Provides decls for the following functions:
//
Expand All @@ -96,8 +125,14 @@ CheckerLibrary ListsCheckerLibrary(int version = kListsExtensionLatestVersion);
//
// <list(T_)>.sort() -> list(T_) where T_ is partially orderable
CompilerLibrary ListsCompilerLibrary(
int version = kListsExtensionLatestVersion);
const ListsExtensionOptions& extension_options = {});

inline CompilerLibrary ListsCompilerLibrary(int version) {
ListsExtensionOptions extension_options;
extension_options.version = version;
return ListsCompilerLibrary(extension_options);
}

} // namespace cel::extensions

#endif // THIRD_PARTY_CEL_CPP_EXTENSIONS_SETS_FUNCTIONS_H_
#endif // THIRD_PARTY_CEL_CPP_EXTENSIONS_LISTS_FUNCTIONS_H_
84 changes: 84 additions & 0 deletions extensions/lists_functions_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
#include "absl/strings/string_view.h"
#include "checker/type_check_issue.h"
#include "checker/validation_result.h"
#include "common/ast.h"
#include "common/source.h"
#include "common/value.h"
#include "common/value_testing.h"
Expand Down Expand Up @@ -123,6 +124,10 @@ INSTANTIATE_TEST_SUITE_P(
// lists.range()
{R"cel(lists.range(4) == [0,1,2,3])cel"},
{R"cel(lists.range(0) == [])cel"},
{R"cel(lists.range(-1))cel",
"lists.range: size must be non-negative, got -1"},
{R"cel(lists.range(1000001))cel",
"lists.range: size 1000001 exceeds maximum allowed (1000000)"},

// .reverse()
{R"cel([5,1,2,3].reverse() == [3,2,1,5])cel"},
Expand Down Expand Up @@ -457,5 +462,84 @@ std::vector<ListsExtensionVersionTestCase> CreateListsExtensionVersionParams() {
INSTANTIATE_TEST_SUITE_P(ListsExtensionVersionTest, ListsExtensionVersionTest,
ValuesIn(CreateListsExtensionVersionParams()));

TEST(ListsFunctionsTest, CustomMaxRangeSizeOption) {
ListsExtensionOptions ext_options;
ext_options.max_range_size = 100;

ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
NewCompilerBuilder(internal::GetTestingDescriptorPool()));
ASSERT_THAT(compiler_builder->AddLibrary(StandardCompilerLibrary()), IsOk());
ASSERT_THAT(compiler_builder->AddLibrary(ListsCompilerLibrary(ext_options)),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, std::move(*compiler_builder).Build());

ASSERT_OK_AND_ASSIGN(ValidationResult result,
compiler->Compile("lists.range(101)", "<input>"));
ASSERT_TRUE(result.IsValid());
ASSERT_OK_AND_ASSIGN(std::unique_ptr<Ast> ast, result.ReleaseAst());

const auto runtime_options = RuntimeOptions{};
ASSERT_OK_AND_ASSIGN(
auto runtime_builder,
CreateStandardRuntimeBuilder(internal::GetTestingDescriptorPool(),
runtime_options));
ASSERT_THAT(RegisterListsFunctions(runtime_builder.function_registry(),
runtime_options, ext_options),
IsOk());
ASSERT_OK_AND_ASSIGN(auto runtime, std::move(runtime_builder).Build());
ASSERT_OK_AND_ASSIGN(auto program, runtime->CreateProgram(std::move(ast)));

google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(Value eval_result,
program->Evaluate(&arena, activation));
EXPECT_THAT(
eval_result,
ErrorValueIs(StatusIs(
testing::_,
HasSubstr("lists.range: size 101 exceeds maximum allowed (100)"))));
}

TEST(ListsFunctionsTest, HardCodedLimitAppliesWhenOptionIsLooser) {
ListsExtensionOptions ext_options;
ext_options.max_range_size = 2000000;

ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
NewCompilerBuilder(internal::GetTestingDescriptorPool()));
ASSERT_THAT(compiler_builder->AddLibrary(StandardCompilerLibrary()), IsOk());
ASSERT_THAT(compiler_builder->AddLibrary(ListsCompilerLibrary(ext_options)),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, std::move(*compiler_builder).Build());

ASSERT_OK_AND_ASSIGN(ValidationResult result,
compiler->Compile("lists.range(1000001)", "<input>"));
ASSERT_TRUE(result.IsValid());
ASSERT_OK_AND_ASSIGN(std::unique_ptr<Ast> ast, result.ReleaseAst());

const auto runtime_options = RuntimeOptions{};
ASSERT_OK_AND_ASSIGN(
auto runtime_builder,
CreateStandardRuntimeBuilder(internal::GetTestingDescriptorPool(),
runtime_options));
ASSERT_THAT(RegisterListsFunctions(runtime_builder.function_registry(),
runtime_options, ext_options),
IsOk());
ASSERT_OK_AND_ASSIGN(auto runtime, std::move(runtime_builder).Build());
ASSERT_OK_AND_ASSIGN(auto program, runtime->CreateProgram(std::move(ast)));

google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(Value eval_result,
program->Evaluate(&arena, activation));
EXPECT_THAT(
eval_result,
ErrorValueIs(StatusIs(
testing::_,
HasSubstr(
"lists.range: size 1000001 exceeds maximum allowed (1000000)"))));
}

} // namespace
} // namespace cel::extensions
Loading