diff --git a/src/include/souffle/RecordTable.h b/src/include/souffle/RecordTable.h index 2a6c1ddcdcc..a93b6c5dce0 100644 --- a/src/include/souffle/RecordTable.h +++ b/src/include/souffle/RecordTable.h @@ -61,4 +61,72 @@ RamDomain pack(RecordTableT&& recordTab, const std::initializer_list& return recordTab.pack(std::data(initlist), initlist.size()); } +/** + * @brief Construct a value belonging to an enumerated algebraic data type. + * + * An ADT is enumerated when none of its branches has fields. + * + * An ADT branch ID is the branch's zero-based position in the lexicographical + * ordering of the ADT's qualified branch names. + */ +inline RamDomain packADTEnum(const RamDomain branch) { + return branch; +} + +/** + * @brief Construct a value belonging to a non-enumerated algebraic data type. + * + * An ADT is non-enumerated when at least one of its branches has a field. This + * function must also be used for an empty branch of such an ADT. + * + * The arguments must already be represented as RamDomain values, using the + * same symbol and record tables as the program that will consume the ADT. This + * helper hides the special record encodings used for empty, unary, and + * multi-argument branches. + * + * @param branch Zero-based position in the lexicographical ordering of the + * ADT's qualified branch names. + * @param arguments Branch field values in declaration order, encoded as + * RamDomain values. + */ +inline RamDomain packADT( + RecordTable& recordTab, const RamDomain branch, const RamDomain* arguments, const std::size_t arity) { + const RamDomain branchValue = arity == 1 ? arguments[0] : recordTab.pack(arguments, arity); + return recordTab.pack({branch, branchValue}); +} + +/** + * @brief Construct a non-enumerated ADT value from an initialization list. + * @param branch Zero-based position in lexicographical branch-name order. + * @param arguments Branch field values in declaration order, encoded as + * RamDomain values. + */ +inline RamDomain packADT( + RecordTable& recordTab, const RamDomain branch, const std::initializer_list& arguments) { + return packADT(recordTab, branch, std::data(arguments), arguments.size()); +} + +/** + * @brief Construct a non-enumerated ADT value from a fixed-size tuple. + * @param branch Zero-based position in lexicographical branch-name order. + * @param arguments Branch field values in declaration order, encoded as + * RamDomain values. + */ +template +RamDomain packADT(RecordTable& recordTab, const RamDomain branch, const Tuple& arguments) { + return packADT(recordTab, branch, arguments.data(), Arity); +} + +/** + * @brief Construct a non-enumerated ADT value from a span. + * @param branch Zero-based position in lexicographical branch-name order. + * @param arguments Branch field values in declaration order, encoded as + * RamDomain values. + */ +template +RamDomain packADT( + RecordTable& recordTab, const RamDomain branch, const span arguments) { + return packADT(recordTab, branch, arguments.data(), arguments.size()); +} + } // namespace souffle diff --git a/src/tests/record_table_test.cpp b/src/tests/record_table_test.cpp index 76e0ad53435..f59e6b6e17b 100644 --- a/src/tests/record_table_test.cpp +++ b/src/tests/record_table_test.cpp @@ -70,6 +70,55 @@ TEST(Pack, InitListHelper) { EXPECT_EQ(3, ptr[2]); } +TEST(PackADT, Enum) { + EXPECT_EQ(2, packADTEnum(2)); +} + +TEST(PackADT, EmptyBranch) { + SpecializedRecordTable<0, 2> recordTable; + + const RamDomain ref = packADT(recordTable, 3, {}); + const RamDomain* adt = recordTable.unpack(ref, 2); + + EXPECT_EQ(3, adt[0]); + EXPECT_EQ(recordTable.pack(nullptr, 0), adt[1]); +} + +TEST(PackADT, UnaryBranch) { + SpecializedRecordTable<2> recordTable; + + const RamDomain ref = packADT(recordTable, 4, {42}); + const RamDomain* adt = recordTable.unpack(ref, 2); + + EXPECT_EQ(4, adt[0]); + EXPECT_EQ(42, adt[1]); +} + +TEST(PackADT, MultiArgumentBranch) { + SpecializedRecordTable<2, 3> recordTable; + const Tuple arguments = {{10, 20, 30}}; + + const RamDomain ref = packADT(recordTable, 5, arguments); + const RamDomain* adt = recordTable.unpack(ref, 2); + const RamDomain* unpackedArguments = recordTable.unpack(adt[1], 3); + + EXPECT_EQ(5, adt[0]); + EXPECT_EQ(10, unpackedArguments[0]); + EXPECT_EQ(20, unpackedArguments[1]); + EXPECT_EQ(30, unpackedArguments[2]); +} + +TEST(PackADT, NestedBranch) { + SpecializedRecordTable<2> recordTable; + + const RamDomain inner = packADT(recordTable, 0, {7}); + const RamDomain outer = packADT(recordTable, 1, {inner}); + const RamDomain* unpackedOuter = recordTable.unpack(outer, 2); + + EXPECT_EQ(1, unpackedOuter[0]); + EXPECT_EQ(inner, unpackedOuter[1]); +} + TEST(Enumerate, Empty) { SpecializedRecordTable<2> recordTable; diff --git a/tests/interface/functors/constructed.csv b/tests/interface/functors/constructed.csv new file mode 100644 index 00000000000..177212a9adf --- /dev/null +++ b/tests/interface/functors/constructed.csv @@ -0,0 +1,3 @@ +"$A(1234)" +"$B(""5678"")" +"$C" diff --git a/tests/interface/functors/functors.cpp b/tests/interface/functors/functors.cpp index 9de7239f413..8d8cffaa50c 100644 --- a/tests/interface/functors/functors.cpp +++ b/tests/interface/functors/functors.cpp @@ -147,4 +147,18 @@ souffle::RamDomain my_to_number_fun( souffle::RamDomain my_identity(souffle::SymbolTable*, souffle::RecordTable*, souffle::RamDomain arg) { return arg; } + +souffle::RamDomain my_make_a( + souffle::SymbolTable*, souffle::RecordTable* recordTable, souffle::RamDomain arg) { + return souffle::packADT(*recordTable, 0, {arg}); +} + +souffle::RamDomain my_make_b( + souffle::SymbolTable*, souffle::RecordTable* recordTable, souffle::RamDomain arg) { + return souffle::packADT(*recordTable, 1, {arg}); +} + +souffle::RamDomain my_make_c(souffle::SymbolTable*, souffle::RecordTable* recordTable) { + return souffle::packADT(*recordTable, 2, {}); +} } // end of extern "C" diff --git a/tests/interface/functors/functors.dl b/tests/interface/functors/functors.dl index 1ad7d1e8e98..99e063cea26 100644 --- a/tests/interface/functors/functors.dl +++ b/tests/interface/functors/functors.dl @@ -93,15 +93,22 @@ L(@myappend(l)) :- L(l), l = [x, _l1], x < 10. .output L // Testing ADTS -.type MyADT = A { x : number } +// Branches are deliberately declared in reverse lexical order. Their IDs are +// still assigned lexically: A = 0, B = 1, and C = 2. +.type MyADT = C {} | B { x : symbol } + | A { x : number } .functor my_to_number_fun(val: MyADT): number stateful .functor my_identity(val: MyADT): MyADT stateful +.functor my_make_a(val: number): MyADT stateful +.functor my_make_b(val: symbol): MyADT stateful +.functor my_make_c(): MyADT stateful .decl gen_adt(val: MyADT) .decl my_to_number(val: number) .decl identity(val: MyADT) +.decl constructed(val: MyADT) gen_adt($A(1234)). gen_adt($B("5678")). @@ -110,4 +117,10 @@ my_to_number(@my_to_number_fun(v)) :- gen_adt(v). identity(@my_identity(v)) :- gen_adt(v). +constructed(@my_make_a(1234)). +constructed(@my_make_b("5678")). +constructed($C()). +constructed(@my_make_c()). + .output my_to_number, identity +.output constructed(rfc4180=true)