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.
- spark_joinery-0.0.1/PKG-INFO +12 -0
- spark_joinery-0.0.1/pyproject.toml +41 -0
- spark_joinery-0.0.1/pyproject.toml.orig +46 -0
- spark_joinery-0.0.1/src/spark_joinery/__init__.py +20 -0
- spark_joinery-0.0.1/src/spark_joinery/collection.py +53 -0
- spark_joinery-0.0.1/src/spark_joinery/dependencies.py +28 -0
- spark_joinery-0.0.1/src/spark_joinery/fixtures.py +23 -0
- spark_joinery-0.0.1/src/spark_joinery/pipeline.py +290 -0
- spark_joinery-0.0.1/src/spark_joinery/schemas.py +188 -0
- spark_joinery-0.0.1/src/spark_joinery/transform.py +212 -0
- spark_joinery-0.0.1/src/spark_joinery/type_inspection.py +124 -0
|
@@ -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}")
|