spark-joinery 0.0.1__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,12 @@
1
+ Metadata-Version: 2.3
2
+ Name: spark-joinery
3
+ Version: 0.0.1
4
+ Summary: Schema safety for pyspark transformations
5
+ Author: Connor Charles
6
+ Author-email: Connor Charles <ccharles.gb@gmail.com>
7
+ Requires-Dist: pyspark>=4.2.0
8
+ Requires-Dist: typer>=0.19.2 ; extra == 'cli'
9
+ Requires-Dist: pydantic>=2.0.0 ; extra == 'pydantic'
10
+ Requires-Python: >=3.13
11
+ Provides-Extra: cli
12
+ Provides-Extra: pydantic
@@ -0,0 +1,41 @@
1
+ [project]
2
+ name = "spark-joinery"
3
+ version = "0.0.1"
4
+ description = "Schema safety for pyspark transformations"
5
+ requires-python = ">=3.13"
6
+ dependencies = ["pyspark>=4.2.0"]
7
+
8
+ [[project.authors]]
9
+ name = "Connor Charles"
10
+ email = "ccharles.gb@gmail.com"
11
+
12
+ [project.optional-dependencies]
13
+ cli = ["typer>=0.19.2"]
14
+ pydantic = ["pydantic>=2.0.0"]
15
+
16
+ [project.scripts]
17
+ spark-joinery = "spark_joinery:main"
18
+
19
+ [tool.pyrefly]
20
+ project-includes = [
21
+ "**/*.py*",
22
+ "**/*.ipynb",
23
+ ]
24
+ search-path = [
25
+ "src",
26
+ "examples",
27
+ ]
28
+
29
+ [build-system]
30
+ requires = ["uv_build>=0.9.27,<0.10.0"]
31
+ build-backend = "uv_build"
32
+
33
+ [dependency-groups]
34
+ dev = [
35
+ "ipykernel>=7.3.0",
36
+ "pydantic>=2.0.0",
37
+ "pyrefly>=1.2.0",
38
+ "pytest>=9.1.0",
39
+ "ruff>=0.15.17",
40
+ "zensical>=0.0.45",
41
+ ]
@@ -0,0 +1,46 @@
1
+ [project]
2
+ name = "spark-joinery"
3
+ version = "0.0.1"
4
+ description = "Schema safety for pyspark transformations"
5
+ authors = [
6
+ { name = "Connor Charles", email = "ccharles.gb@gmail.com" }
7
+ ]
8
+ requires-python = ">=3.13"
9
+ dependencies = [
10
+ "pyspark>=4.2.0",
11
+ ]
12
+
13
+ [tool.pyrefly]
14
+ project-includes = [
15
+ "**/*.py*",
16
+ "**/*.ipynb",
17
+ ]
18
+ search-path = [
19
+ "src",
20
+ "examples",
21
+ ]
22
+
23
+ [project.optional-dependencies]
24
+ cli = [
25
+ "typer>=0.19.2",
26
+ ]
27
+ pydantic = [
28
+ "pydantic>=2.0.0",
29
+ ]
30
+
31
+ [project.scripts]
32
+ spark-joinery = "spark_joinery:main"
33
+
34
+ [build-system]
35
+ requires = ["uv_build>=0.9.27,<0.10.0"]
36
+ build-backend = "uv_build"
37
+
38
+ [dependency-groups]
39
+ dev = [
40
+ "ipykernel>=7.3.0",
41
+ "pydantic>=2.0.0",
42
+ "pyrefly>=1.2.0",
43
+ "pytest>=9.1.0",
44
+ "ruff>=0.15.17",
45
+ "zensical>=0.0.45",
46
+ ]
@@ -0,0 +1,20 @@
1
+ from .collection import Collection
2
+ from .dependencies import Context, PipelineContext
3
+ from .pipeline import (
4
+ ExecutablePipeline,
5
+ Pipeline,
6
+ PipelineExecutionError,
7
+ Step,
8
+ )
9
+ from .transform import transform
10
+
11
+ __all__ = [
12
+ "Collection",
13
+ "Context",
14
+ "ExecutablePipeline",
15
+ "Pipeline",
16
+ "PipelineExecutionError",
17
+ "PipelineContext",
18
+ "Step",
19
+ "transform",
20
+ ]
@@ -0,0 +1,53 @@
1
+ from typing import Any, Callable, ParamSpec, TypeVar, overload
2
+
3
+ from .schemas import CoercionMode
4
+ from .transform import TransformSpec, _inspect_transform, _wrap_transform
5
+
6
+ P = ParamSpec("P")
7
+ R = TypeVar("R")
8
+
9
+
10
+ class Collection:
11
+ def __init__(self) -> None:
12
+ self._specs: dict[Callable[..., Any], TransformSpec] = {}
13
+
14
+ @overload
15
+ def transform(
16
+ self,
17
+ f: Callable[P, R],
18
+ *,
19
+ validate_input: CoercionMode | None = "project_all",
20
+ validate_output: CoercionMode | None = "project_all",
21
+ ) -> Callable[P, R]: ...
22
+
23
+ @overload
24
+ def transform(
25
+ self,
26
+ f: None = None,
27
+ *,
28
+ validate_input: CoercionMode | None = "project_all",
29
+ validate_output: CoercionMode | None = "project_all",
30
+ ) -> Callable[[Callable[P, R]], Callable[P, R]]: ...
31
+
32
+ def transform(
33
+ self,
34
+ f: Callable[P, R] | None = None,
35
+ *,
36
+ validate_input: CoercionMode | None = "project_all",
37
+ validate_output: CoercionMode | None = "project_all",
38
+ ):
39
+ def decorator(fn: Callable[P, R]) -> Callable[P, R]:
40
+ spec = _inspect_transform(fn)
41
+ wrapper = _wrap_transform(
42
+ fn,
43
+ spec,
44
+ validate_input=validate_input,
45
+ validate_output=validate_output,
46
+ )
47
+ self._specs[wrapper] = spec
48
+ return wrapper
49
+
50
+ if f is None:
51
+ return decorator
52
+
53
+ return decorator(f)
@@ -0,0 +1,28 @@
1
+ from typing import Any, Sequence
2
+
3
+
4
+ class Context:
5
+ """Marks a transform parameter as resolved from a PipelineContext."""
6
+
7
+ def __init__(self, type_: type | None = None) -> None:
8
+ self.type = type_
9
+
10
+
11
+ class PipelineContext:
12
+ """Registry of dependency values, looked up by type."""
13
+
14
+ def __init__(self, *, values: Sequence[Any] = ()) -> None:
15
+ self._by_type: dict[type, Any] = {}
16
+ for value in values:
17
+ value_type = type(value)
18
+ if value_type in self._by_type:
19
+ raise ValueError(
20
+ f"duplicate dependency type {value_type!r} in PipelineContext"
21
+ )
22
+ self._by_type[value_type] = value
23
+
24
+ def resolve(self, param_type: type, marker: "Context") -> Any:
25
+ lookup_type = marker.type or param_type
26
+ if lookup_type not in self._by_type:
27
+ raise KeyError(f"no dependency registered for type {lookup_type!r}")
28
+ return self._by_type[lookup_type]
@@ -0,0 +1,23 @@
1
+ from collections.abc import Sequence
2
+ from typing import TypeVar
3
+ from pyspark.sql import DataFrame, SparkSession
4
+ from spark_joinery import schemas
5
+
6
+ T = TypeVar("T")
7
+
8
+
9
+ def get_dataframe(
10
+ spark: SparkSession, row_type: type[T], rows: Sequence[T]
11
+ ) -> DataFrame:
12
+ for index, row in enumerate(rows):
13
+ if not isinstance(row, row_type):
14
+ raise ValueError(
15
+ f"Row {index} of type {row.__class__.__name__}. Expected type {row_type.__name__}"
16
+ )
17
+ schema = schemas.get_spark_schema_from_model(row_type)
18
+
19
+ serialized_rows = [
20
+ row.model_dump(mode="python") if hasattr(row, "model_dump") else row
21
+ for row in rows
22
+ ]
23
+ return spark.createDataFrame(serialized_rows, schema)
@@ -0,0 +1,290 @@
1
+ from __future__ import annotations
2
+
3
+ import inspect
4
+ from dataclasses import dataclass
5
+ from typing import Any, Callable, Sequence
6
+
7
+ from pyspark.sql import DataFrame, SparkSession, types
8
+
9
+ from .collection import Collection
10
+ from .dependencies import Context, PipelineContext
11
+ from .schemas import CoercionMode
12
+
13
+ Transform = Callable[..., DataFrame | None]
14
+
15
+
16
+ class PipelineExecutionError(RuntimeError):
17
+ pass
18
+
19
+
20
+ @dataclass(eq=False)
21
+ class Step:
22
+ name: str
23
+ transform: Transform
24
+ _pipeline: Pipeline
25
+ _input_schemas: dict[str, types.StructType]
26
+ _output_schema: types.StructType | None
27
+ _spark_parameter: str | None
28
+ _context_parameters: dict[str, tuple[type, Context]]
29
+ _upstream_steps: list[Step]
30
+ _explicit_bindings: dict[Step, str]
31
+
32
+
33
+ @dataclass(frozen=True)
34
+ class _ExecutionStep:
35
+ step: Step
36
+ dataframe_bindings: tuple[tuple[str, Step], ...]
37
+
38
+
39
+ class ExecutablePipeline:
40
+ def __init__(self, execution_steps: Sequence[_ExecutionStep]):
41
+ self._execution_steps = tuple(execution_steps)
42
+
43
+ def run(
44
+ self, spark: SparkSession, context: PipelineContext | None = None
45
+ ) -> dict[str, DataFrame]:
46
+ if not isinstance(spark, SparkSession):
47
+ raise TypeError("run() requires a SparkSession")
48
+
49
+ outputs: dict[str, DataFrame] = {}
50
+ for execution_step in self._execution_steps:
51
+ step = execution_step.step
52
+ arguments: dict[str, Any] = {
53
+ parameter_name: outputs[upstream_step.name]
54
+ for parameter_name, upstream_step in execution_step.dataframe_bindings
55
+ }
56
+ if step._spark_parameter is not None:
57
+ arguments[step._spark_parameter] = spark
58
+
59
+ for parameter_name, (
60
+ param_type,
61
+ marker,
62
+ ) in step._context_parameters.items():
63
+ if context is None:
64
+ raise PipelineExecutionError(
65
+ f"step '{step.name}' requires a PipelineContext but none was provided"
66
+ )
67
+ try:
68
+ arguments[parameter_name] = context.resolve(param_type, marker)
69
+ except KeyError as error:
70
+ raise PipelineExecutionError(
71
+ f"step '{step.name}' failed to resolve dependency"
72
+ ) from error
73
+
74
+ try:
75
+ result = step.transform(**arguments)
76
+ if step._output_schema is not None and not isinstance(
77
+ result, DataFrame
78
+ ):
79
+ raise TypeError(
80
+ f"step '{step.name}' returned a non-DataFrame value"
81
+ )
82
+ except Exception as error:
83
+ if isinstance(error, PipelineExecutionError):
84
+ raise
85
+ raise PipelineExecutionError(
86
+ f"Pipeline step '{step.name}' failed"
87
+ ) from error
88
+
89
+ if isinstance(result, DataFrame):
90
+ outputs[step.name] = result
91
+
92
+ return outputs
93
+
94
+
95
+ class Pipeline:
96
+ def __init__(self, collections: Sequence[Collection] = ()):
97
+ self._collection = Collection()
98
+ self._collections: tuple[Collection, ...] = (self._collection, *collections)
99
+ self._steps: dict[str, Step] = {}
100
+ self._validated = False
101
+ self._executable: ExecutablePipeline | None = None
102
+
103
+ def transform(
104
+ self,
105
+ f=None,
106
+ *,
107
+ validate_input: CoercionMode | None = "project_all",
108
+ validate_output: CoercionMode | None = "project_all",
109
+ ):
110
+ return self._collection.transform(
111
+ f,
112
+ validate_input=validate_input,
113
+ validate_output=validate_output,
114
+ )
115
+
116
+ def _resolve_spec(self, transform: Transform):
117
+ for collection in self._collections:
118
+ spec = collection._specs.get(transform)
119
+ if spec is not None:
120
+ return spec
121
+ name = getattr(transform, "__name__", repr(transform))
122
+ raise TypeError(f"'{name}' is not registered in this pipeline's collections")
123
+
124
+ def add_step(self, transform: Transform, name: str) -> Step:
125
+ self._ensure_mutable()
126
+ if name in self._steps:
127
+ raise ValueError(f"step name '{name}' is already registered")
128
+
129
+ spec = self._resolve_spec(transform)
130
+ input_schemas = spec.input_schemas
131
+ output_schema = spec.output_schema
132
+
133
+ spark_parameter = spec.spark_parameter
134
+ parameter_names = set(inspect.signature(transform).parameters)
135
+ supported_parameters = (
136
+ set(input_schemas)
137
+ | ({spark_parameter} if spark_parameter is not None else set())
138
+ | set(spec.context_parameters)
139
+ )
140
+ unsupported_parameters = parameter_names - supported_parameters
141
+ if unsupported_parameters:
142
+ unsupported = sorted(unsupported_parameters)[0]
143
+ raise TypeError(
144
+ "pipeline steps only support SparkSession, annotated DataFrame, "
145
+ f"and Context-annotated parameters; unsupported parameter '{unsupported}'"
146
+ )
147
+
148
+ if not input_schemas and spark_parameter is None:
149
+ raise TypeError("source step must declare a SparkSession")
150
+
151
+ step = Step(
152
+ name=name,
153
+ transform=transform,
154
+ _pipeline=self,
155
+ _input_schemas=input_schemas,
156
+ _output_schema=output_schema,
157
+ _spark_parameter=spark_parameter,
158
+ _context_parameters=spec.context_parameters,
159
+ _upstream_steps=[],
160
+ _explicit_bindings={},
161
+ )
162
+ self._steps[name] = step
163
+ return step
164
+
165
+ def connect(
166
+ self, upstream: Step, downstream: Step, *, param: str | None = None
167
+ ) -> None:
168
+ self._ensure_mutable()
169
+ self._ensure_owned(upstream)
170
+ self._ensure_owned(downstream)
171
+ if upstream not in downstream._upstream_steps:
172
+ downstream._upstream_steps.append(upstream)
173
+ if param is not None:
174
+ for other_upstream, other_param in downstream._explicit_bindings.items():
175
+ if other_param == param and other_upstream is not upstream:
176
+ raise ValueError(
177
+ f"step '{downstream.name}' parameter '{param}' is already "
178
+ f"bound to upstream step '{other_upstream.name}'"
179
+ )
180
+ downstream._explicit_bindings[upstream] = param
181
+
182
+ def connect_many(
183
+ self,
184
+ upstream_steps: Sequence[Step],
185
+ downstream: Step,
186
+ ) -> None:
187
+ self._ensure_mutable()
188
+ if not upstream_steps:
189
+ raise ValueError("connect_many requires at least one upstream step")
190
+ for upstream in upstream_steps:
191
+ self.connect(upstream, downstream)
192
+
193
+ def validate(self) -> ExecutablePipeline:
194
+ if self._executable is not None:
195
+ return self._executable
196
+
197
+ execution_steps: list[_ExecutionStep] = []
198
+ for step in self._topological_order():
199
+ bindings: list[tuple[str, Step]] = []
200
+ matched_upstream: set[Step] = set()
201
+ for parameter_name, expected_schema in step._input_schemas.items():
202
+ explicit_upstream = next(
203
+ (
204
+ upstream
205
+ for upstream, bound_param in step._explicit_bindings.items()
206
+ if bound_param == parameter_name
207
+ ),
208
+ None,
209
+ )
210
+ if explicit_upstream is not None:
211
+ if explicit_upstream._output_schema != expected_schema:
212
+ raise ValueError(
213
+ f"step '{step.name}' parameter '{parameter_name}' is "
214
+ f"explicitly bound to '{explicit_upstream.name}' but its "
215
+ "output schema does not match"
216
+ )
217
+ bindings.append((parameter_name, explicit_upstream))
218
+ matched_upstream.add(explicit_upstream)
219
+ continue
220
+
221
+ candidates = [
222
+ upstream
223
+ for upstream in step._upstream_steps
224
+ if upstream not in step._explicit_bindings
225
+ ]
226
+ matches = [
227
+ upstream
228
+ for upstream in candidates
229
+ if upstream._output_schema == expected_schema
230
+ ]
231
+ if not matches:
232
+ raise ValueError(
233
+ f"step '{step.name}' parameter '{parameter_name}' has no "
234
+ "upstream step provides a matching schema"
235
+ )
236
+ if len(matches) > 1:
237
+ names = ", ".join(upstream.name for upstream in matches)
238
+ raise ValueError(
239
+ f"step '{step.name}' has ambiguous upstream schema for "
240
+ f"parameter '{parameter_name}': {names}"
241
+ )
242
+ upstream = matches[0]
243
+ bindings.append((parameter_name, upstream))
244
+ matched_upstream.add(upstream)
245
+
246
+ unmatched = [
247
+ upstream
248
+ for upstream in step._upstream_steps
249
+ if upstream not in matched_upstream
250
+ ]
251
+ if unmatched:
252
+ names = ", ".join(upstream.name for upstream in unmatched)
253
+ raise ValueError(
254
+ f"step '{step.name}' has upstream output that does not match "
255
+ f"a parameter: {names}"
256
+ )
257
+
258
+ execution_steps.append(_ExecutionStep(step, tuple(bindings)))
259
+
260
+ self._validated = True
261
+ self._executable = ExecutablePipeline(execution_steps)
262
+ return self._executable
263
+
264
+ def _ensure_mutable(self) -> None:
265
+ if self._validated:
266
+ raise RuntimeError("pipeline has already been validated")
267
+
268
+ def _ensure_owned(self, step: Step) -> None:
269
+ if step._pipeline is not self:
270
+ raise ValueError("steps must belong to the same pipeline")
271
+
272
+ def _topological_order(self) -> list[Step]:
273
+ states: dict[Step, int] = {}
274
+ ordered: list[Step] = []
275
+
276
+ def visit(step: Step) -> None:
277
+ state = states.get(step, 0)
278
+ if state == 1:
279
+ raise ValueError("pipeline contains a cycle")
280
+ if state == 2:
281
+ return
282
+ states[step] = 1
283
+ for upstream in step._upstream_steps:
284
+ visit(upstream)
285
+ states[step] = 2
286
+ ordered.append(step)
287
+
288
+ for step in self._steps.values():
289
+ visit(step)
290
+ return ordered
@@ -0,0 +1,188 @@
1
+ from dataclasses import fields, is_dataclass
2
+ from functools import partial
3
+ from typing import Any, Callable, TypeVar, Literal
4
+
5
+ from pyspark.errors import PySparkAssertionError
6
+ from pyspark.sql import Column, DataFrame, functions as F, types
7
+ from pyspark.testing import assertSchemaEqual
8
+
9
+ from . import type_inspection
10
+
11
+ T = TypeVar("T")
12
+
13
+ CoercionMode = Literal["coerce", "project_all", "project", "strict", "strict_null"]
14
+
15
+
16
+ def is_pydantic_model(klass: Any) -> bool:
17
+ # Avoid hard dependency on pydantic; detect by BaseModel class attributes.
18
+ return (
19
+ isinstance(klass, type)
20
+ and hasattr(klass, "model_fields")
21
+ and isinstance(getattr(klass, "model_fields"), dict)
22
+ )
23
+
24
+
25
+ def is_schema_model(klass: Any) -> bool:
26
+ return is_dataclass(klass) or is_pydantic_model(klass)
27
+
28
+
29
+ def _get_model_fields(klass: type[Any]) -> list[tuple[str, Any]]:
30
+ if is_dataclass(klass):
31
+ return [(field.name, field.type) for field in fields(klass)]
32
+
33
+ if is_pydantic_model(klass):
34
+ model_fields = getattr(klass, "model_fields")
35
+ return [
36
+ (name, field_info.annotation)
37
+ for name, field_info in model_fields.items()
38
+ if field_info.annotation is not None
39
+ ]
40
+
41
+ raise ValueError(f"{klass.__name__} is neither a dataclass nor a pydantic model")
42
+
43
+
44
+ def _get_spark_field_type(field_type: Any) -> types.DataType:
45
+ normalized_type = type_inspection.normalize_python_type(field_type)
46
+
47
+ if is_schema_model(normalized_type):
48
+ return get_spark_schema_from_model(normalized_type)
49
+
50
+ if type_inspection.is_list(normalized_type):
51
+ element_type = type_inspection.get_list_element_type(normalized_type)
52
+ element_spark_type = _get_spark_field_type(element_type)
53
+ return types.ArrayType(element_spark_type, True)
54
+
55
+ return type_inspection.get_spark_type_from_python_type(normalized_type)
56
+
57
+
58
+ def get_spark_schema_from_model(klass: type[T]) -> types.StructType:
59
+ if not is_schema_model(klass):
60
+ raise ValueError(
61
+ f"{klass.__name__} is neither a dataclass nor a pydantic model"
62
+ )
63
+
64
+ struct_fields = [
65
+ types.StructField(name, _get_spark_field_type(field_type), True)
66
+ for name, field_type in _get_model_fields(klass)
67
+ ]
68
+
69
+ return types.StructType(struct_fields)
70
+
71
+
72
+ def get_spark_schema_from_dataclass(klass: type[T]) -> types.StructType:
73
+ if not is_dataclass(klass):
74
+ raise ValueError(f"{klass.__name__} is not a dataclass")
75
+
76
+ return get_spark_schema_from_model(klass)
77
+
78
+
79
+ def dataframe_is_schema(dataframe: DataFrame, klass: type[T]) -> bool:
80
+ expected_schema = get_spark_schema_from_dataclass(klass)
81
+ return dataframe.schema == expected_schema
82
+
83
+
84
+ def dataframe_is_model_schema(dataframe: DataFrame, klass: type[T]) -> bool:
85
+ expected_schema = get_spark_schema_from_model(klass)
86
+ return dataframe.schema == expected_schema
87
+
88
+
89
+ def _coerce_strict(
90
+ dataframe: DataFrame, schema: types.StructType, *, ignore_nullable: bool
91
+ ) -> DataFrame:
92
+ try:
93
+ assertSchemaEqual(dataframe.schema, schema, ignoreNullable=ignore_nullable)
94
+ except PySparkAssertionError as e:
95
+ raise ValueError("Schema mismatch") from e
96
+ return dataframe
97
+
98
+
99
+ def _build_projected_column(
100
+ column: Column,
101
+ source_type: types.DataType,
102
+ target_type: types.DataType,
103
+ *,
104
+ cast: bool,
105
+ recurse: bool,
106
+ ) -> Column:
107
+ if recurse and isinstance(target_type, types.StructType):
108
+ if not isinstance(source_type, types.StructType):
109
+ raise ValueError(f"Expected struct type but found {source_type}")
110
+
111
+ source_fields = {field.name: field.dataType for field in source_type.fields}
112
+ nested_columns = []
113
+ for target_field in target_type.fields:
114
+ if target_field.name not in source_fields:
115
+ raise ValueError(f"Missing field '{target_field.name}'")
116
+ nested_column = _build_projected_column(
117
+ column.getField(target_field.name),
118
+ source_fields[target_field.name],
119
+ target_field.dataType,
120
+ cast=cast,
121
+ recurse=recurse,
122
+ )
123
+ nested_columns.append(nested_column.alias(target_field.name))
124
+ return F.struct(*nested_columns)
125
+
126
+ if recurse and isinstance(target_type, types.ArrayType):
127
+ if not isinstance(source_type, types.ArrayType):
128
+ raise ValueError(f"Expected array type but found {source_type}")
129
+
130
+ return F.transform(
131
+ column,
132
+ lambda element: _build_projected_column(
133
+ element,
134
+ source_type.elementType,
135
+ target_type.elementType,
136
+ cast=cast,
137
+ recurse=recurse,
138
+ ),
139
+ )
140
+
141
+ if source_type == target_type:
142
+ return column
143
+
144
+ if cast:
145
+ return column.cast(target_type)
146
+
147
+ raise ValueError(f"Type mismatch: expected {target_type}, found {source_type}")
148
+
149
+
150
+ def _project_fields(
151
+ dataframe: DataFrame, schema: types.StructType, *, cast: bool, recurse: bool
152
+ ) -> DataFrame:
153
+ source_fields = {field.name: field.dataType for field in dataframe.schema.fields}
154
+
155
+ missing_fields = [
156
+ field.name for field in schema.fields if field.name not in source_fields
157
+ ]
158
+ if missing_fields:
159
+ raise ValueError(f"Missing columns: {missing_fields}")
160
+
161
+ columns = [
162
+ _build_projected_column(
163
+ F.col(field.name),
164
+ source_fields[field.name],
165
+ field.dataType,
166
+ cast=cast,
167
+ recurse=recurse,
168
+ ).alias(field.name)
169
+ for field in schema.fields
170
+ ]
171
+ return dataframe.select(*columns)
172
+
173
+
174
+ _MODE_HANDLERS: dict[
175
+ CoercionMode, Callable[[DataFrame, types.StructType], DataFrame]
176
+ ] = {
177
+ "strict": partial(_coerce_strict, ignore_nullable=True),
178
+ "strict_null": partial(_coerce_strict, ignore_nullable=False),
179
+ "project": partial(_project_fields, cast=False, recurse=False),
180
+ "project_all": partial(_project_fields, cast=False, recurse=True),
181
+ "coerce": partial(_project_fields, cast=True, recurse=True),
182
+ }
183
+
184
+
185
+ def coerce_dataframe(
186
+ dataframe: DataFrame, schema: types.StructType, mode: CoercionMode
187
+ ) -> DataFrame:
188
+ return _MODE_HANDLERS[mode](dataframe, schema)
@@ -0,0 +1,212 @@
1
+ from dataclasses import dataclass
2
+ from functools import wraps
3
+ import inspect
4
+ from typing import (
5
+ Annotated,
6
+ Any,
7
+ Callable,
8
+ ParamSpec,
9
+ TypeVar,
10
+ cast,
11
+ get_args,
12
+ get_origin,
13
+ get_type_hints,
14
+ overload,
15
+ )
16
+
17
+ from pyspark.sql import DataFrame, SparkSession
18
+ from pyspark.sql import types
19
+
20
+ from . import schemas
21
+ from .dependencies import Context
22
+ from .schemas import CoercionMode
23
+
24
+ P = ParamSpec("P")
25
+ R = TypeVar("R")
26
+
27
+
28
+ def _get_annotated_dataframe_schema_model(annotation: Any) -> type[Any] | None:
29
+ if get_origin(annotation) is not Annotated:
30
+ return None
31
+
32
+ annotated_args = get_args(annotation)
33
+ if len(annotated_args) < 2:
34
+ return None
35
+
36
+ base_type = annotated_args[0]
37
+ metadata = annotated_args[1:]
38
+ if base_type is not DataFrame:
39
+ return None
40
+
41
+ for metadata_value in metadata:
42
+ if isinstance(metadata_value, type) and schemas.is_schema_model(metadata_value):
43
+ return metadata_value
44
+
45
+ return None
46
+
47
+
48
+ def _get_context_marker(annotation: Any) -> tuple[type, Context] | None:
49
+ if get_origin(annotation) is not Annotated:
50
+ return None
51
+
52
+ annotated_args = get_args(annotation)
53
+ if len(annotated_args) < 2:
54
+ return None
55
+
56
+ base_type = annotated_args[0]
57
+ metadata = annotated_args[1:]
58
+ for metadata_value in metadata:
59
+ if isinstance(metadata_value, Context):
60
+ return base_type, metadata_value
61
+
62
+ return None
63
+
64
+
65
+ @dataclass(frozen=True)
66
+ class TransformSpec:
67
+ input_schemas: dict[str, types.StructType]
68
+ output_schema: types.StructType | None
69
+ spark_parameter: str | None
70
+ context_parameters: dict[str, tuple[type, Context]]
71
+
72
+
73
+ def _get_spark_parameter(f: Any) -> str | None:
74
+ type_hints = get_type_hints(f)
75
+ spark_parameters = [
76
+ parameter_name
77
+ for parameter_name in inspect.signature(f).parameters
78
+ if type_hints.get(parameter_name) is SparkSession
79
+ ]
80
+ if len(spark_parameters) > 1:
81
+ raise TypeError("transform may declare only one SparkSession parameter")
82
+ return spark_parameters[0] if spark_parameters else None
83
+
84
+
85
+ def _inspect_transform(f: Any) -> TransformSpec:
86
+ signature = inspect.signature(f)
87
+ type_hints = get_type_hints(f, include_extras=True)
88
+ input_schemas: dict[str, types.StructType] = {}
89
+ context_parameters: dict[str, tuple[type, Context]] = {}
90
+
91
+ for parameter_name in signature.parameters:
92
+ parameter_type = type_hints.get(parameter_name)
93
+ if parameter_type is None:
94
+ continue
95
+
96
+ dataframe_schema_model = _get_annotated_dataframe_schema_model(parameter_type)
97
+ if dataframe_schema_model is not None:
98
+ input_schemas[parameter_name] = schemas.get_spark_schema_from_model(
99
+ dataframe_schema_model
100
+ )
101
+ continue
102
+
103
+ context_marker = _get_context_marker(parameter_type)
104
+ if context_marker is not None:
105
+ context_parameters[parameter_name] = context_marker
106
+
107
+ output_schema = None
108
+ return_type = type_hints.get("return")
109
+ if return_type is not None:
110
+ dataframe_schema_model = _get_annotated_dataframe_schema_model(return_type)
111
+ if dataframe_schema_model is not None:
112
+ output_schema = schemas.get_spark_schema_from_model(dataframe_schema_model)
113
+
114
+ return TransformSpec(
115
+ input_schemas=input_schemas,
116
+ output_schema=output_schema,
117
+ spark_parameter=_get_spark_parameter(f),
118
+ context_parameters=context_parameters,
119
+ )
120
+
121
+
122
+ def _wrap_transform(
123
+ fn: Callable[P, R],
124
+ spec: TransformSpec,
125
+ *,
126
+ validate_input: CoercionMode | None,
127
+ validate_output: CoercionMode | None,
128
+ ) -> Callable[P, R]:
129
+ signature = inspect.signature(fn)
130
+
131
+ @wraps(fn)
132
+ def wrapper(*args: P.args, **kwds: P.kwargs) -> R:
133
+ bound_arguments = signature.bind(*args, **kwds)
134
+ bound_arguments.apply_defaults()
135
+
136
+ if validate_input is not None:
137
+ for parameter_name, expected_schema in spec.input_schemas.items():
138
+ value = bound_arguments.arguments.get(parameter_name)
139
+ if not isinstance(value, DataFrame):
140
+ raise TypeError(
141
+ f"Parameter '{parameter_name}' must be a pyspark.sql.DataFrame"
142
+ )
143
+
144
+ try:
145
+ bound_arguments.arguments[parameter_name] = (
146
+ schemas.coerce_dataframe(value, expected_schema, validate_input)
147
+ )
148
+ except ValueError as e:
149
+ raise ValueError(
150
+ f"Schema mismatch for parameter '{parameter_name}'"
151
+ ) from e
152
+
153
+ result = fn(*bound_arguments.args, **bound_arguments.kwargs)
154
+
155
+ if validate_output is not None and spec.output_schema is not None:
156
+ if not isinstance(result, DataFrame):
157
+ raise TypeError(
158
+ f"Return value from '{fn.__name__}' must be a pyspark.sql.DataFrame"
159
+ )
160
+
161
+ try:
162
+ result = cast(
163
+ R,
164
+ schemas.coerce_dataframe(
165
+ result, spec.output_schema, validate_output
166
+ ),
167
+ )
168
+ except ValueError as e:
169
+ raise ValueError(f"Return schema mismatch for '{fn.__name__}'") from e
170
+
171
+ return result
172
+
173
+ return wrapper
174
+
175
+
176
+ @overload
177
+ def transform(
178
+ f: Callable[P, R],
179
+ *,
180
+ validate_input: CoercionMode | None = "project_all",
181
+ validate_output: CoercionMode | None = "project_all",
182
+ ) -> Callable[P, R]: ...
183
+
184
+
185
+ @overload
186
+ def transform(
187
+ f: None = None,
188
+ *,
189
+ validate_input: CoercionMode | None = "project_all",
190
+ validate_output: CoercionMode | None = "project_all",
191
+ ) -> Callable[[Callable[P, R]], Callable[P, R]]: ...
192
+
193
+
194
+ def transform(
195
+ f: Callable[P, R] | None = None,
196
+ *,
197
+ validate_input: CoercionMode | None = "project_all",
198
+ validate_output: CoercionMode | None = "project_all",
199
+ ):
200
+ def decorator(fn: Callable[P, R]) -> Callable[P, R]:
201
+ spec = _inspect_transform(fn)
202
+ return _wrap_transform(
203
+ fn,
204
+ spec,
205
+ validate_input=validate_input,
206
+ validate_output=validate_output,
207
+ )
208
+
209
+ if f is None:
210
+ return decorator
211
+
212
+ return decorator(f)
@@ -0,0 +1,124 @@
1
+ import datetime
2
+ import decimal
3
+ from types import UnionType
4
+
5
+ from pyspark.sql import types
6
+ from typing import Any, Annotated, Literal, get_origin, get_args, Union
7
+
8
+ NoneType = type(None)
9
+
10
+
11
+ SPARK_MAX_DECIMAL_PRECISION = 38
12
+ DEFAULT_FRACTIONAL_DIGITS = 18
13
+
14
+
15
+ def allows_none(tp: Any) -> bool:
16
+ origin = get_origin(tp)
17
+ args = get_args(tp)
18
+
19
+ # Optional[T] is Union[T, None]
20
+ if origin is Union and NoneType in args:
21
+ return True
22
+
23
+ # PEP 604 syntax: T | None
24
+ if origin is UnionType and NoneType in args:
25
+ return True
26
+
27
+ return False
28
+
29
+
30
+ def is_list(tp: Any) -> bool:
31
+ tp = normalize_python_type(tp)
32
+ origin = get_origin(tp)
33
+ return origin is list
34
+
35
+
36
+ def get_list_element_type(tp: Any) -> Any:
37
+ tp = normalize_python_type(tp)
38
+ if not is_list(tp):
39
+ raise ValueError(f"{tp} is not a list type")
40
+ args = get_args(tp)
41
+ if len(args) != 1:
42
+ raise ValueError(f"List type {tp} should have exactly one type argument")
43
+ return args[0]
44
+
45
+
46
+ def normalize_python_type(python_type: Any) -> Any:
47
+ origin = get_origin(python_type)
48
+ args = get_args(python_type)
49
+
50
+ if origin is Annotated and args:
51
+ return normalize_python_type(args[0])
52
+
53
+ if origin in (Union, UnionType) and args:
54
+ non_none_args = [arg for arg in args if arg is not NoneType]
55
+ if len(non_none_args) == 1 and len(non_none_args) != len(args):
56
+ return normalize_python_type(non_none_args[0])
57
+
58
+ return python_type
59
+
60
+
61
+ def spark_type_from_annotated(python_type: Any) -> types.DataType | None:
62
+ origin = get_origin(python_type)
63
+ args = get_args(python_type)
64
+
65
+ if origin is not Annotated or len(args) < 2:
66
+ return None
67
+
68
+ for metadata in args[1:]:
69
+ if isinstance(metadata, types.DataType):
70
+ return metadata
71
+
72
+ return None
73
+
74
+
75
+ def spark_type_from_literal(python_type: Any) -> types.DataType | None:
76
+ if get_origin(python_type) is not Literal:
77
+ return None
78
+
79
+ literal_values = get_args(python_type)
80
+ if not literal_values:
81
+ raise ValueError("Literal type must include at least one value")
82
+
83
+ spark_types = {
84
+ get_spark_type_from_python_type(type(literal_value))
85
+ for literal_value in literal_values
86
+ }
87
+
88
+ if len(spark_types) != 1:
89
+ raise ValueError(
90
+ f"Literal values must resolve to a single Spark type: {python_type}"
91
+ )
92
+
93
+ return next(iter(spark_types))
94
+
95
+
96
+ def get_spark_type_from_python_type(python_type: type) -> types.DataType:
97
+ annotated_spark_type = spark_type_from_annotated(python_type)
98
+ if annotated_spark_type is not None:
99
+ return annotated_spark_type
100
+
101
+ literal_spark_type = spark_type_from_literal(python_type)
102
+ if literal_spark_type is not None:
103
+ return literal_spark_type
104
+
105
+ python_type = normalize_python_type(python_type)
106
+
107
+ if python_type is int:
108
+ return types.IntegerType()
109
+ elif python_type is str:
110
+ return types.StringType()
111
+ elif python_type is float:
112
+ return types.FloatType()
113
+ elif python_type is bool:
114
+ return types.BooleanType()
115
+ elif python_type is datetime.datetime:
116
+ return types.TimestampType()
117
+ elif python_type is datetime.date:
118
+ return types.DateType()
119
+ elif python_type is decimal.Decimal:
120
+ return types.DecimalType(SPARK_MAX_DECIMAL_PRECISION, DEFAULT_FRACTIONAL_DIGITS)
121
+ elif python_type is bytes:
122
+ return types.BinaryType()
123
+ else:
124
+ raise ValueError(f"Unsupported Python type: {python_type}")