|
15 | 15 | ) |
16 | 16 | from podio_gen.podio_config_reader import PodioConfigReader |
17 | 17 | from podio_gen.generator_base import ClassGeneratorBaseMixin, write_file_if_changed |
18 | | -from podio_gen.generator_utils import DataType, DataModelJSONEncoder |
| 18 | +from podio_gen.generator_utils import DataType, DataModelJSONEncoder, DefinitionError |
19 | 19 |
|
20 | 20 | REPORT_TEXT = """ |
21 | 21 | PODIO Data Model |
|
24 | 24 | Read instructions in the README.md to run your first example! |
25 | 25 | """ |
26 | 26 |
|
| 27 | +ARROW_PRIMITIVE_TYPES = { |
| 28 | + "bool": "arrow::boolean()", |
| 29 | + "char": "arrow::int8()", |
| 30 | + "short": "arrow::int16()", |
| 31 | + "int": "arrow::int32()", |
| 32 | + "long": "arrow::int64()", |
| 33 | + "long long": "arrow::int64()", |
| 34 | + "unsigned": "arrow::uint32()", |
| 35 | + "unsigned int": "arrow::uint32()", |
| 36 | + "unsigned long": "arrow::uint64()", |
| 37 | + "unsigned long long": "arrow::uint64()", |
| 38 | + "float": "arrow::float32()", |
| 39 | + "double": "arrow::float64()", |
| 40 | + "int16_t": "arrow::int16()", |
| 41 | + "int32_t": "arrow::int32()", |
| 42 | + "int64_t": "arrow::int64()", |
| 43 | + "uint16_t": "arrow::uint16()", |
| 44 | + "uint32_t": "arrow::uint32()", |
| 45 | + "uint64_t": "arrow::uint64()", |
| 46 | + "std::int16_t": "arrow::int16()", |
| 47 | + "std::int32_t": "arrow::int32()", |
| 48 | + "std::int64_t": "arrow::int64()", |
| 49 | + "std::uint16_t": "arrow::uint16()", |
| 50 | + "std::uint32_t": "arrow::uint32()", |
| 51 | + "std::uint64_t": "arrow::uint64()", |
| 52 | + "std::string": "arrow::utf8()", |
| 53 | +} |
| 54 | + |
27 | 55 |
|
28 | 56 | class IncludeFrom(IntEnum): |
29 | 57 | """Enum to signify if an include is needed and from where it should come""" |
@@ -85,6 +113,9 @@ def post_process(self, datamodel): |
85 | 113 | if "ROOT" in self.io_handlers: |
86 | 114 | self._create_selection_xml() |
87 | 115 |
|
| 116 | + if "ARROW" in self.io_handlers: |
| 117 | + self._write_arrow_mapper_header(datamodel) |
| 118 | + |
88 | 119 | if the_links := datamodel["links"]: |
89 | 120 | self._write_links_registration_file(the_links) |
90 | 121 | self._write_all_collections_header() |
@@ -122,6 +153,76 @@ def do_process_datatype(self, name, datatype): |
122 | 153 |
|
123 | 154 | return datatype |
124 | 155 |
|
| 156 | + def _write_arrow_mapper_header(self, datamodel): |
| 157 | + """A generated helper that exposes the datamodel as an Arrow schema""" |
| 158 | + datatypes = [] |
| 159 | + for datatype in datamodel["datatypes"]: |
| 160 | + datatype["arrow_fields"] = self._arrow_fields(datatype) |
| 161 | + datatypes.append(datatype) |
| 162 | + |
| 163 | + data = { |
| 164 | + "package_name": self.package_name, |
| 165 | + "schema_version": self.datamodel.schema_version, |
| 166 | + "datatypes": datatypes, |
| 167 | + } |
| 168 | + self._write_file( |
| 169 | + "ArrowMapper.h", |
| 170 | + self._eval_template("ArrowMapper.h.jinja2", data), |
| 171 | + ) |
| 172 | + |
| 173 | + def _arrow_fields(self, datatype): |
| 174 | + """Create Arrow field expressions for the members and relations of a datatype""" |
| 175 | + fields = [] |
| 176 | + fields.extend( |
| 177 | + self._arrow_field(member.name, self._arrow_type(member)) |
| 178 | + for member in datatype["Members"] |
| 179 | + ) |
| 180 | + fields.extend( |
| 181 | + self._arrow_field(member.name, f"arrow::list({self._arrow_type(member)})") |
| 182 | + for member in datatype["VectorMembers"] |
| 183 | + ) |
| 184 | + fields.extend( |
| 185 | + self._arrow_field(relation.name, "objectRefType()") |
| 186 | + for relation in datatype["OneToOneRelations"] |
| 187 | + ) |
| 188 | + fields.extend( |
| 189 | + self._arrow_field(relation.name, "arrow::list(objectRefType())") |
| 190 | + for relation in datatype["OneToManyRelations"] |
| 191 | + ) |
| 192 | + return fields |
| 193 | + |
| 194 | + def _arrow_field(self, name, type_expr, nullable=False): |
| 195 | + """Create a C++ arrow::field expression""" |
| 196 | + nullable_arg = ", true" if nullable else "" |
| 197 | + return f'arrow::field("{name}", {type_expr}{nullable_arg})' |
| 198 | + |
| 199 | + def _arrow_type(self, member): |
| 200 | + """Map a parsed podio member to an Arrow C++ DataType expression""" |
| 201 | + if member.is_array: |
| 202 | + value_type = self._arrow_type_from_name(member.array_type) |
| 203 | + return f"arrow::fixed_size_list({value_type}, {member.array_size})" |
| 204 | + |
| 205 | + return self._arrow_type_from_name(member.full_type) |
| 206 | + |
| 207 | + def _arrow_type_from_name(self, type_name): |
| 208 | + """Map a C++ type name from the datamodel to an Arrow C++ DataType expression""" |
| 209 | + type_name = type_name.removeprefix("::") |
| 210 | + if type_name in ARROW_PRIMITIVE_TYPES: |
| 211 | + return ARROW_PRIMITIVE_TYPES[type_name] |
| 212 | + |
| 213 | + if type_name in self.datamodel.components: |
| 214 | + return self._arrow_struct_type(self.datamodel.components[type_name]["Members"]) |
| 215 | + |
| 216 | + if self.upstream_edm and type_name in self.upstream_edm.components: |
| 217 | + return self._arrow_struct_type(self.upstream_edm.components[type_name]["Members"]) |
| 218 | + |
| 219 | + raise DefinitionError(f"Cannot map '{type_name}' to an Arrow type") |
| 220 | + |
| 221 | + def _arrow_struct_type(self, members): |
| 222 | + """Create an Arrow struct expression for a component definition""" |
| 223 | + fields = [self._arrow_field(member.name, self._arrow_type(member)) for member in members] |
| 224 | + return "arrow::struct_({" + ", ".join(fields) + "})" |
| 225 | + |
125 | 226 | def do_process_interface(self, _, interface): |
126 | 227 | """Process an interface definition and generate the necessary code""" |
127 | 228 | interface["include_types"] = [ |
|
0 commit comments