diff --git a/spark_auto_mapper/automappers/complex.py b/spark_auto_mapper/automappers/complex.py index e4b8b4bb..7200ce74 100644 --- a/spark_auto_mapper/automappers/complex.py +++ b/spark_auto_mapper/automappers/complex.py @@ -52,7 +52,7 @@ def __init__( column_name: str mapper: AutoMapperDataTypeBase for column_name, mapper in entity.get_child_mappers().items(): - if column_name == "extension": + if not enable_schema_pruning and column_name == "extension": extension_schema: Union[StructType, DataType, None] # since there is a column called extension then get the schema with extension extension_schema = mapper.get_schema( diff --git a/spark_auto_mapper/automappers/with_column_base.py b/spark_auto_mapper/automappers/with_column_base.py index ea3841bf..db87684d 100644 --- a/spark_auto_mapper/automappers/with_column_base.py +++ b/spark_auto_mapper/automappers/with_column_base.py @@ -11,6 +11,7 @@ from spark_auto_mapper.automappers.automapper_base import AutoMapperBase from spark_auto_mapper.automappers.check_schema_result import CheckSchemaResult from spark_auto_mapper.data_types.data_type_base import AutoMapperDataTypeBase +from spark_auto_mapper.schema_pruning.schema_pruner import SchemaPruner from spark_auto_mapper.type_definitions.defined_types import AutoMapperAnyDataType from spark_auto_mapper.helpers.value_parser import AutoMapperValueParser @@ -53,11 +54,23 @@ def get_column_spec(self, source_df: Optional[DataFrame]) -> Column: child: AutoMapperDataTypeBase = self.value if self.column_schema: if self.enable_schema_pruning: - self.value.set_schema( + self.value.mark_used_fields_in_schema( column_name=self.dst_column, column_path=self.dst_column, + field=self.column_schema, column_data_type=self.column_schema.dataType, ) + # remove all fields without "used" property + SchemaPruner.prune_schema( + field=self.column_schema, + field_data_type=self.column_schema.dataType, + ) + + # self.value.set_schema( + # column_name=self.dst_column, + # column_path=self.dst_column, + # column_data_type=self.column_schema.dataType, + # ) column_spec = child.get_column_spec( source_df=source_df, current_column=None, parent_columns=None ) @@ -76,16 +89,16 @@ def get_column_spec(self, source_df: Optional[DataFrame]) -> Column: # if the type has a schema then apply it if self.column_schema: column_data_type: DataType = self.column_schema.dataType - if self.enable_schema_pruning: - # first disable generation of null properties since we are doing schema reduction - self.value.include_null_properties(include_null_properties=False) - # second ask the mapper to reduce schema that is not used - column_data_type = self.value.filter_schema_by_fields_present( - column_name=self.dst_column, - column_path=self.dst_column, - column_data_type=column_data_type, - skip_null_properties=True, - ) + # if self.enable_schema_pruning: + # # first disable generation of null properties since we are doing schema reduction + # self.value.include_null_properties(include_null_properties=False) + # # second ask the mapper to reduce schema that is not used + # column_data_type = self.value.filter_schema_by_fields_present( + # column_name=self.dst_column, + # column_path=self.dst_column, + # column_data_type=column_data_type, + # skip_null_properties=True, + # ) column_spec = column_spec.cast(column_data_type) # if dst_column already exists in source_df then prepend with ___ to make it unique if source_df is not None and self.dst_column in source_df.columns: @@ -143,7 +156,8 @@ def check_schema( desired_schema=desired_schema, ) return CheckSchemaResult(result=result) - except AnalysisException: + except AnalysisException as e: + print(e) return None else: return None diff --git a/spark_auto_mapper/data_types/data_type_base.py b/spark_auto_mapper/data_types/data_type_base.py index fda9e389..1adc536f 100644 --- a/spark_auto_mapper/data_types/data_type_base.py +++ b/spark_auto_mapper/data_types/data_type_base.py @@ -1024,6 +1024,152 @@ def set_schema( column_data_type=column_data_type, ) + def _mark_used_fields_in_schema_for_array( + self, + *, + column_name: Optional[str], + column_path: Optional[str], + field: StructField, + column_data_type: DataType, + ) -> None: + assert isinstance(field, StructField) + assert isinstance(column_data_type, DataType) + assert isinstance( + column_data_type, ArrayType + ), f"{type(column_data_type)} should be ArrayType for {column_name} with path {column_path}" + + element_type = column_data_type.elementType + # self.set_children_schema(element_type) + children: Union[ + AutoMapperDataTypeBase, List[AutoMapperDataTypeBase] + ] = self.children + assert isinstance(children, list), f"{type(children)} should be a list" + if len(children) > 0: + child: AutoMapperDataTypeBase + for child in children: + child.mark_used_fields_in_schema( + column_name=column_name, + column_path=column_path, + field=field, + column_data_type=element_type, + ) + + # noinspection PyUnusedLocal + def _mark_used_fields_in_schema_for_struct( + self, + *, + column_name: Optional[str], + column_path: Optional[str], + field: StructField, + column_data_type: DataType, + ) -> None: + assert isinstance(field, StructField) + assert isinstance(column_data_type, DataType) + + assert isinstance( + column_data_type, StructType + ), f"{type(column_data_type)} should be StructType for {column_name} with path {column_path}" + + children: Union[ + "AutoMapperDataTypeBase", List["AutoMapperDataTypeBase"] + ] = self.children + if isinstance(children, list) and len(children) > 0: + child: "AutoMapperDataTypeBase" + for index, child in enumerate(children): + assert isinstance(child, AutoMapperDataTypeBase), f"{type(child)}" + if not child.column_name: + continue + assert child.column_name, f"No column name for {child}" + clean_child_name: str = PythonKeywordCleaner.from_python_safe( + child.column_name + ) + matching_fields = [ + f for f in column_data_type.fields if f.name == clean_child_name + ] + if len(matching_fields) == 0: + pass + assert len(matching_fields) == 1, ( + f"Schema match failed for column {column_path}.{clean_child_name}" + f" in schema fields" + f": [{','.join([f.name for f in column_data_type.fields])}]" + ) + child_field: StructField = matching_fields[0] + child.mark_used_fields_in_schema( + column_name=child_field.name, + column_path=f"{column_path}.{child_field.name}", + field=child_field, + column_data_type=child_field.dataType, + ) + elif not isinstance(children, list) and children is not None: + child = children + assert child.column_name + clean_child_name = PythonKeywordCleaner.from_python_safe(child.column_name) + matching_fields = [ + f for f in column_data_type.fields if f.name == clean_child_name + ] + assert len(matching_fields) == 1, ( + f"Schema match failed for column {column_path}.{clean_child_name}" + f" in schema fields" + f": [{','.join([f.name for f in column_data_type.fields])}]" + ) + child_field = matching_fields[0] + child.mark_used_fields_in_schema( + column_name=child_field.name, + column_path=f"{column_path}.{child_field.name}", + field=child_field, + column_data_type=child_field.dataType, + ) + + def mark_used_fields_in_schema( + self, + *, + column_name: Optional[str], + column_path: Optional[str], + field: StructField, + column_data_type: DataType, + ) -> None: + """ + Sets the fields as used schema for this AutoMapper type + + + :param column_name: column name + :param column_path: full path to column + :param field: schema field for this mapper + :param column_data_type: schema for this mapper + """ + + assert isinstance(field, StructField) + assert isinstance(column_data_type, DataType) + + # self.schema = column_data_type + + # if this is a basic type so nothing to do + if not isinstance(column_data_type, StructType) and not isinstance( + column_data_type, ArrayType + ): + setattr(field, "used", True) + return + + if isinstance(column_data_type, ArrayType): + self._mark_used_fields_in_schema_for_array( + column_name=column_name, + column_path=column_path, + field=field, + column_data_type=column_data_type, + ) + return + + assert isinstance( + column_data_type, StructType + ), f"{type(column_data_type)} should be StructType for {column_name} with path {column_path}" + + self._mark_used_fields_in_schema_for_struct( + column_name=column_name, + column_path=column_path, + field=field, + column_data_type=column_data_type, + ) + @property @abstractmethod def children( diff --git a/spark_auto_mapper/schema_pruning/__init__.py b/spark_auto_mapper/schema_pruning/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/spark_auto_mapper/schema_pruning/annotated_struct_type.py b/spark_auto_mapper/schema_pruning/annotated_struct_type.py new file mode 100644 index 00000000..4c063873 --- /dev/null +++ b/spark_auto_mapper/schema_pruning/annotated_struct_type.py @@ -0,0 +1,12 @@ +from typing import List + +from pyspark.sql.types import StructType, StructField + + +class AnnotatedStructType: + def __init__(self, fields: List[StructField]) -> None: + pass + + @staticmethod + def parse(struct_type: StructType) -> "AnnotatedStructType": + pass diff --git a/spark_auto_mapper/schema_pruning/schema_pruner.py b/spark_auto_mapper/schema_pruning/schema_pruner.py new file mode 100644 index 00000000..1e84f2cd --- /dev/null +++ b/spark_auto_mapper/schema_pruning/schema_pruner.py @@ -0,0 +1,41 @@ +from pyspark.sql.types import StructType, StructField, DataType, ArrayType + + +class SchemaPruner: + @staticmethod + def _prune_schema_array(*, field: StructField, field_data_type: DataType) -> None: + assert isinstance(field_data_type, ArrayType) + + SchemaPruner.prune_schema( + field=field, field_data_type=field_data_type.elementType + ) + + @staticmethod + def _prune_schema_struct(*, field: StructField, field_data_type: DataType) -> None: + assert isinstance(field_data_type, StructType) + + # remove any fields that don't have the "used" field + field_data_type.fields = [ + f for f in field_data_type.fields if hasattr(f, "used") + ] + field_data_type.names = [ + n + for n in field_data_type.names + if n in [f.name for f in field_data_type.fields] + ] + # remove the used tag + for f in field_data_type.fields: + SchemaPruner.prune_schema(field=f, field_data_type=field_data_type) + delattr(f, "used") + + @staticmethod + def prune_schema(*, field: StructField, field_data_type: DataType) -> None: + if isinstance(field.dataType, StructType): + SchemaPruner._prune_schema_struct( + field=field, field_data_type=field_data_type + ) + + if isinstance(field.dataType, ArrayType): + SchemaPruner._prune_schema_array( + field=field, field_data_type=field_data_type + ) diff --git a/tests/schema_pruner/__init__.py b/tests/schema_pruner/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/schema_pruner/test_schema_pruner.py b/tests/schema_pruner/test_schema_pruner.py new file mode 100644 index 00000000..86b0ebba --- /dev/null +++ b/tests/schema_pruner/test_schema_pruner.py @@ -0,0 +1,17 @@ +from pyspark.sql.types import StructField, StructType, StringType, IntegerType + +from spark_auto_mapper.schema_pruning.schema_pruner import SchemaPruner + + +def test_schema_pruner() -> None: + field: StructField = StructField( + "foo", + StructType( + [ + StructField("prop1", StringType()), + StructField("prop2", IntegerType()), + ] + ), + ) + + SchemaPruner.prune_schema(field=field, field_data_type=field.dataType) diff --git a/tests/schema_pruning/test_automapper_schema_pruning_with_defined_class.py b/tests/schema_pruning/test_automapper_schema_pruning_with_defined_class.py index 4fda1666..8005686c 100644 --- a/tests/schema_pruning/test_automapper_schema_pruning_with_defined_class.py +++ b/tests/schema_pruning/test_automapper_schema_pruning_with_defined_class.py @@ -53,8 +53,7 @@ def test_auto_mapper_schema_pruning_with_defined_class( # Act mapper = AutoMapper( - view="members", - source_view="patients", + view="members", source_view="patients", enable_schema_pruning=True ).complex(MyClass(name=A.column("last_name"), age=A.number(A.column("my_age")))) assert isinstance(mapper, AutoMapper) diff --git a/tests/schema_pruning/test_automapper_schema_pruning_with_extension_nested_children.py b/tests/schema_pruning/test_automapper_schema_pruning_with_extension_nested_children.py new file mode 100644 index 00000000..0caeff53 --- /dev/null +++ b/tests/schema_pruning/test_automapper_schema_pruning_with_extension_nested_children.py @@ -0,0 +1,347 @@ +from typing import Dict, Optional, Union, List + +# noinspection PyPackageRequirements +from pyspark.sql import SparkSession, Column, DataFrame + +# noinspection PyPackageRequirements +from pyspark.sql.functions import col + +# noinspection PyPackageRequirements +from pyspark.sql.types import ( + ArrayType, + LongType, + StringType, + StructField, + StructType, + TimestampType, + DataType, +) +from spark_data_frame_comparer.schema_comparer import ( + SchemaComparer, + SchemaComparerResult, +) + +from spark_auto_mapper.automappers.automapper import AutoMapper +from spark_auto_mapper.data_types.complex.complex_base import ( + AutoMapperDataTypeComplexBase, +) +from spark_auto_mapper.data_types.data_type_base import AutoMapperDataTypeBase +from spark_auto_mapper.data_types.list import AutoMapperList +from spark_auto_mapper.data_types.number import AutoMapperNumberDataType +from spark_auto_mapper.data_types.text_like_base import AutoMapperTextLikeBase +from spark_auto_mapper.helpers.automapper_helpers import AutoMapperHelpers as A +from spark_auto_mapper.type_definitions.defined_types import AutoMapperDateInputType +from tests.conftest import clean_spark_session + + +class MyProcessingStatusExtensionItem(AutoMapperDataTypeComplexBase): + # noinspection PyPep8Naming + def __init__( + self, + url: str, + valueString: Optional[AutoMapperTextLikeBase] = None, + ) -> None: + super().__init__( + url=url, + valueString=valueString, + ) + + +class MyProcessingStatusExtensionItem2(AutoMapperDataTypeComplexBase): + # noinspection PyPep8Naming + def __init__( + self, + url: str, + valueString: Optional[AutoMapperTextLikeBase] = None, + valueDateTime: Optional[AutoMapperDateInputType] = None, + ) -> None: + super().__init__( + url=url, + valueString=valueString, + valueDateTime=valueDateTime, + ) + + +def get_extension_schema( + include_extension: bool = False, +) -> StructType: + return StructType( + [ + StructField("url", StringType()), + StructField("extra", StringType(), True), + StructField( + "extension", + ArrayType( + StructType( + [ + StructField("url", StringType()), + StructField("valueString", StringType()), + StructField("valueUrl", StringType()), + StructField("valueDateTime", TimestampType()), + ] + ) + ), + ), + ] + ) + + +class MyProcessingStatusExtension(AutoMapperDataTypeComplexBase): + # noinspection PyPep8Naming + def __init__( + self, + processing_status: AutoMapperTextLikeBase, + request_id: AutoMapperTextLikeBase, + ) -> None: + definition_base_url = "https://raw.githubusercontent.com/imranq2/SparkAutoMapper.FHIR/main/StructureDefinition/" + processing_status_extensions = [ + MyProcessingStatusExtensionItem( + url="processing_status", + valueString=processing_status, + ), + MyProcessingStatusExtensionItem( + url="request_id", + valueString=request_id, + ), + ] + self.extensions = processing_status_extensions + super().__init__( + url=definition_base_url, + extension=AutoMapperList(processing_status_extensions), + ) + + def include_null_properties(self, include_null_properties: bool) -> None: + for item in self.extensions: + item.include_null_properties( + include_null_properties=include_null_properties + ) + + def get_schema( + self, include_extension: bool, extension_fields: Optional[List[str]] = None + ) -> Optional[Union[StructType, DataType]]: + return StructType( + [ + StructField("url", StringType()), + StructField("extra", StringType(), True), + StructField( + "extension", + ArrayType( + StructType( + [ + StructField("url", StringType()), + StructField("valueString", StringType()), + ] + ) + ), + ), + ] + ) + + def get_value( + self, + value: AutoMapperDataTypeBase, + source_df: Optional[DataFrame], + current_column: Optional[Column], + ) -> Column: + return super().get_value(value, source_df, current_column) + + +class MyProcessingStatusExtension2(AutoMapperDataTypeComplexBase): + # noinspection PyPep8Naming + def __init__( + self, + processing_status: AutoMapperTextLikeBase, + request_id: AutoMapperTextLikeBase, + date_processed: Optional[AutoMapperDateInputType] = None, + ) -> None: + definition_base_url = "https://raw.githubusercontent.com/imranq2/SparkAutoMapper.FHIR/main/StructureDefinition/" + processing_status_extensions = [ + MyProcessingStatusExtensionItem2( + url="processing_status", + valueString=processing_status, + ), + MyProcessingStatusExtensionItem2( + url="request_id", + valueString=request_id, + ), + MyProcessingStatusExtensionItem2( + url="date_processed", + valueDateTime=date_processed, + ), + ] + self.extensions = processing_status_extensions + super().__init__( + url=definition_base_url, + extension=AutoMapperList(processing_status_extensions), + ) + + def include_null_properties(self, include_null_properties: bool) -> None: + for item in self.extensions: + item.include_null_properties( + include_null_properties=include_null_properties + ) + + def get_schema( + self, include_extension: bool, extension_fields: Optional[List[str]] = None + ) -> Optional[Union[StructType, DataType]]: + return StructType( + [ + StructField("url", StringType()), + StructField("extra", StringType(), True), + StructField( + "extension", + ArrayType( + StructType( + [ + StructField("url", StringType()), + StructField("valueString", StringType()), + StructField("valueDateTime", TimestampType()), + ] + ) + ), + ), + ] + ) + + def get_value( + self, + value: AutoMapperDataTypeBase, + source_df: Optional[DataFrame], + current_column: Optional[Column], + ) -> Column: + return super().get_value(value, source_df, current_column) + + +class MyClass(AutoMapperDataTypeComplexBase): + def __init__( + self, + name: AutoMapperTextLikeBase, + age: AutoMapperNumberDataType, + extension: AutoMapperList[AutoMapperDataTypeComplexBase], + ) -> None: + super().__init__(name=name, age=age, extension=extension) + + def get_schema( + self, include_extension: bool, extension_fields: Optional[List[str]] = None + ) -> Optional[Union[StructType, DataType]]: + schema: StructType = StructType( + [ + StructField("name", StringType(), False), + StructField("extra", LongType(), True), + StructField("age", LongType(), True), + StructField( + "extension", + ArrayType(get_extension_schema()), + True, + ), + ] + ) + return schema + + +def test_auto_mapper_schema_pruning_with_extension_nested_children( + spark_session: SparkSession, +) -> None: + # Arrange + clean_spark_session(spark_session) + + spark_session.createDataFrame( + [ + (1, "Qureshi", "Imran", 45), + (2, "Vidal", "Michael", 35), + ], + ["member_id", "last_name", "first_name", "my_age"], + ).createOrReplaceTempView("patients") + + source_df: DataFrame = spark_session.table("patients") + + # Act + mapper = AutoMapper( + view="members", + source_view="patients", + enable_schema_pruning=True, + skip_schema_validation=[], + ).complex( + MyClass( + name=A.column("last_name"), + age=A.number(A.column("my_age")), + extension=AutoMapperList( + [ + MyProcessingStatusExtension( + processing_status=A.text("foo"), + request_id=A.text("bar"), + ), + MyProcessingStatusExtension2( + processing_status=A.text("foo"), + request_id=A.text("bar"), + date_processed=A.date("2021-01-01"), + ), + ] + ), + ) + ) + + # schema = get_extension_schema() + + # annotated_struct_type: AnnotatedStructType = AnnotatedStructType.parse(struct_type=schema) + + assert isinstance(mapper, AutoMapper) + sql_expressions: Dict[str, Column] = mapper.get_column_specs(source_df=source_df) + for column_name, sql_expression in sql_expressions.items(): + print(f"{column_name}: {sql_expression}") + + result_df: DataFrame = mapper.transform(df=source_df) + + # Assert + assert str(sql_expressions["name"]) == str( + col("b.last_name").cast("string").alias("name") + ) + assert str(sql_expressions["age"]) == str(col("b.my_age").cast("long").alias("age")) + + result_df.printSchema() + result_df.show(truncate=False) + + assert result_df.where("member_id == 1").select("name").collect()[0][0] == "Qureshi" + + assert dict(result_df.dtypes)["age"] in ("int", "long", "bigint") + + # confirm schema + expected_schema: StructType = StructType( + [ + StructField("name", StringType(), False), + StructField("age", LongType(), True), + StructField( + "extension", + ArrayType( + StructType( + [ + StructField("url", StringType()), + StructField( + "extension", + ArrayType( + StructType( + [ + StructField("url", StringType()), + StructField("valueString", StringType()), + StructField( + "valueDateTime", TimestampType() + ), + ] + ) + ), + ), + ] + ) + ), + True, + ), + ] + ) + + result: SchemaComparerResult = SchemaComparer.compare_schema( + parent_column_name=None, + source_schema=result_df.schema, + desired_schema=expected_schema, + ) + + assert result.errors == [], str(result)