Samara 0.2__tar.gz → 0.3__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.
Files changed (78) hide show
  1. {samara-0.2 → samara-0.3}/PKG-INFO +2 -1
  2. {samara-0.2 → samara-0.3}/pyproject.toml +2 -1
  3. {samara-0.2 → samara-0.3}/src/samara/__init__.py +0 -1
  4. {samara-0.2 → samara-0.3}/src/samara/exceptions.py +0 -3
  5. {samara-0.2 → samara-0.3}/src/samara/utils/file.py +2 -36
  6. {samara-0.2 → samara-0.3}/src/samara/utils/logger.py +20 -8
  7. {samara-0.2 → samara-0.3}/src/samara/workflow/controller.py +0 -1
  8. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/model_extract.py +2 -2
  9. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/model_load.py +4 -2
  10. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/model_transform.py +0 -1
  11. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/extract.py +21 -48
  12. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/job.py +9 -10
  13. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/load.py +22 -49
  14. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/schema.py +19 -14
  15. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/session.py +54 -5
  16. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transform.py +13 -40
  17. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/base.py +8 -8
  18. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/filter.py +1 -7
  19. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/join.py +4 -23
  20. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/select.py +0 -5
  21. {samara-0.2 → samara-0.3}/LICENSE +0 -0
  22. {samara-0.2 → samara-0.3}/README.md +0 -0
  23. {samara-0.2 → samara-0.3}/src/samara/__main__.py +0 -0
  24. {samara-0.2 → samara-0.3}/src/samara/alert/__init__.py +0 -0
  25. {samara-0.2 → samara-0.3}/src/samara/alert/channels/__init__.py +0 -0
  26. {samara-0.2 → samara-0.3}/src/samara/alert/channels/base.py +0 -0
  27. {samara-0.2 → samara-0.3}/src/samara/alert/channels/email.py +0 -0
  28. {samara-0.2 → samara-0.3}/src/samara/alert/channels/file.py +0 -0
  29. {samara-0.2 → samara-0.3}/src/samara/alert/channels/http.py +0 -0
  30. {samara-0.2 → samara-0.3}/src/samara/alert/controller.py +0 -0
  31. {samara-0.2 → samara-0.3}/src/samara/alert/rules/__init__.py +0 -0
  32. {samara-0.2 → samara-0.3}/src/samara/alert/rules/base.py +0 -0
  33. {samara-0.2 → samara-0.3}/src/samara/alert/rules/env_vars_matches.py +0 -0
  34. {samara-0.2 → samara-0.3}/src/samara/alert/rules/exception_regex.py +0 -0
  35. {samara-0.2 → samara-0.3}/src/samara/alert/template.py +0 -0
  36. {samara-0.2 → samara-0.3}/src/samara/alert/trigger.py +0 -0
  37. {samara-0.2 → samara-0.3}/src/samara/cli.py +0 -0
  38. {samara-0.2 → samara-0.3}/src/samara/settings.py +0 -0
  39. {samara-0.2 → samara-0.3}/src/samara/telemetry.py +0 -0
  40. {samara-0.2 → samara-0.3}/src/samara/types.py +0 -0
  41. {samara-0.2 → samara-0.3}/src/samara/utils/__init__.py +0 -0
  42. {samara-0.2 → samara-0.3}/src/samara/utils/http.py +0 -0
  43. {samara-0.2 → samara-0.3}/src/samara/workflow/__init__.py +0 -0
  44. {samara-0.2 → samara-0.3}/src/samara/workflow/actions/__init__.py +0 -0
  45. {samara-0.2 → samara-0.3}/src/samara/workflow/actions/base.py +0 -0
  46. {samara-0.2 → samara-0.3}/src/samara/workflow/actions/http.py +0 -0
  47. {samara-0.2 → samara-0.3}/src/samara/workflow/actions/move_or_copy_job_files.py +0 -0
  48. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/__init__.py +0 -0
  49. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/hooks.py +0 -0
  50. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/__init__.py +0 -0
  51. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/model_job.py +0 -0
  52. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/__init__.py +0 -0
  53. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_aggregate.py +0 -0
  54. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_cast.py +0 -0
  55. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_distinct.py +0 -0
  56. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_drop.py +0 -0
  57. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_dropduplicates.py +0 -0
  58. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_dropna.py +0 -0
  59. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_filter.py +0 -0
  60. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_groupby.py +0 -0
  61. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_join.py +0 -0
  62. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_orderby.py +0 -0
  63. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_pivot.py +0 -0
  64. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_select.py +0 -0
  65. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/models/transforms/model_withcolumn.py +0 -0
  66. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/polars/.gitkeep +0 -0
  67. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/__init__.py +0 -0
  68. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/__init__.py +0 -0
  69. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/aggregate.py +0 -0
  70. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/cast.py +0 -0
  71. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/distinct.py +0 -0
  72. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/drop.py +0 -0
  73. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/dropduplicates.py +0 -0
  74. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/dropna.py +0 -0
  75. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/groupby.py +0 -0
  76. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/orderby.py +0 -0
  77. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/pivot.py +0 -0
  78. {samara-0.2 → samara-0.3}/src/samara/workflow/jobs/spark/transforms/withcolumn.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: Samara
3
- Version: 0.2
3
+ Version: 0.3
4
4
  Summary: Config Driven ETL Framework
5
5
  License-File: LICENSE
6
6
  Author: Krijn van der Burg
@@ -22,6 +22,7 @@ Requires-Dist: pyjson5 (>=1.6.9,<2.0.0)
22
22
  Requires-Dist: pyspark (>=4.0.1,<5.0.0)
23
23
  Requires-Dist: pyyaml (>=6.0.1,<7.0.0)
24
24
  Requires-Dist: requests (>=2.32.5,<3.0.0)
25
+ Requires-Dist: rich (>=14.3.3,<15.0.0)
25
26
  Requires-Dist: structlog (>=25.4.0,<26.0.0)
26
27
  Description-Content-Type: text/markdown
27
28
 
@@ -1,6 +1,6 @@
1
1
  [tool.poetry]
2
2
  name = "Samara"
3
- version = "0.2"
3
+ version = "0.3"
4
4
  description = "Config Driven ETL Framework"
5
5
  authors = ["Krijn van der Burg"]
6
6
  readme = "README.md"
@@ -24,6 +24,7 @@ click = "^8.1.0"
24
24
  opentelemetry-api = "^1.37.0"
25
25
  opentelemetry-sdk = "^1.37.0"
26
26
  opentelemetry-exporter-otlp = "^1.37.0"
27
+ rich = "^14.3.3"
27
28
 
28
29
  [tool.poetry.group.test.dependencies]
29
30
  pre_commit = "*"
@@ -7,7 +7,6 @@ engines and extensible components.
7
7
 
8
8
  Key capabilities:
9
9
  - Define pipelines via configuration (sources, transforms, destinations)
10
- - Multi-engine architecture (Pandas, Polars, and more)
11
10
  - Configurable alert system with multiple notification channels
12
11
  - Event-triggered custom actions at pipeline stages
13
12
  - Engine-agnostic configuration supporting different backends
@@ -14,9 +14,6 @@ meaningful error states to the operating system.
14
14
  """
15
15
 
16
16
  import enum
17
- from typing import TypeVar
18
-
19
- K = TypeVar("K") # Key type
20
17
 
21
18
 
22
19
  class ExitCode(enum.IntEnum):
@@ -51,9 +51,7 @@ class FileHandler(ABC):
51
51
  File access is deferred until read operations are performed.
52
52
  Validation occurs at read time, not during initialization.
53
53
  """
54
- logger.debug("Initializing %s for path: %s", self.__class__.__name__, str(filepath))
55
54
  self.filepath = filepath
56
- logger.debug("%s initialized successfully for: %s", self.__class__.__name__, str(filepath))
57
55
 
58
56
  def _file_exists(self) -> None:
59
57
  """Verify the file exists.
@@ -62,11 +60,9 @@ class FileHandler(ABC):
62
60
  FileNotFoundError: If the file does not exist.
63
61
  OSError: If a system-level error occurs while checking existence.
64
62
  """
65
- logger.debug("Checking file existence: %s", str(self.filepath))
66
63
  if not self.filepath.exists():
67
64
  logger.error("File not found: %s", str(self.filepath))
68
65
  raise FileNotFoundError(f"File not found: {self.filepath}")
69
- logger.debug("File exists: %s", str(self.filepath))
70
66
 
71
67
  def _is_file(self) -> None:
72
68
  """Verify the path points to a regular file (not a directory).
@@ -75,11 +71,9 @@ class FileHandler(ABC):
75
71
  IsADirectoryError: If the path is a directory.
76
72
  OSError: If the path is not a regular file or a system-level error occurs.
77
73
  """
78
- logger.debug("Checking if path is a regular file: %s", str(self.filepath))
79
74
  if not self.filepath.is_file():
80
75
  logger.error("Path is not a regular file: %s", str(self.filepath))
81
76
  raise OSError(f"Expected a file but found directory or invalid path: '{self.filepath}'")
82
- logger.debug("Path is a regular file: %s", str(self.filepath))
83
77
 
84
78
  def _read_permission(self) -> None:
85
79
  """Verify the file is readable by the current process.
@@ -88,11 +82,9 @@ class FileHandler(ABC):
88
82
  PermissionError: If the file is not readable.
89
83
  OSError: If a system-level error occurs while checking permissions.
90
84
  """
91
- logger.debug("Checking read permissions for file: %s", str(self.filepath))
92
85
  if not os.access(self.filepath, os.R_OK):
93
86
  logger.error("Read permission denied for file: %s", str(self.filepath))
94
87
  raise PermissionError(f"Permission denied: Cannot read file '{self.filepath}'")
95
- logger.debug("Read permissions validated for file: %s", str(self.filepath))
96
88
 
97
89
  def _file_not_empty(self) -> None:
98
90
  """Verify the file contains data (not empty).
@@ -100,12 +92,10 @@ class FileHandler(ABC):
100
92
  Raises:
101
93
  OSError: If the file is empty or a system-level error occurs accessing metadata.
102
94
  """
103
- logger.debug("Checking if file is empty: %s", str(self.filepath))
104
95
  file_size = self.filepath.stat().st_size
105
96
  if file_size == 0:
106
97
  logger.error("File is empty: %s", str(self.filepath))
107
98
  raise OSError(f"File is empty: {self.filepath}")
108
- logger.debug("File not empty: %s (size: %d bytes)", str(self.filepath), file_size)
109
99
 
110
100
  def _file_size_limits(self, max_size: int = DEFAULT_MAX_SIZE) -> None:
111
101
  """Verify file size is within specified limits.
@@ -116,7 +106,6 @@ class FileHandler(ABC):
116
106
  Raises:
117
107
  OSError: If the file exceeds size limits or a system-level error occurs.
118
108
  """
119
- logger.debug("Checking file size limits for: %s (max allowed: %d bytes)", str(self.filepath), max_size)
120
109
  file_size = self.filepath.stat().st_size
121
110
 
122
111
  if file_size > max_size:
@@ -128,10 +117,6 @@ class FileHandler(ABC):
128
117
  )
129
118
  raise OSError(f"File too large: '{self.filepath}' ({file_size:,} bytes exceeds {max_size:,} bytes limit)")
130
119
 
131
- logger.debug(
132
- "File size within limits: %s (size: %d bytes, max: %d bytes)", str(self.filepath), file_size, max_size
133
- )
134
-
135
120
  def _text_file(self) -> None:
136
121
  """Verify the file contains readable text (not binary data).
137
122
 
@@ -139,7 +124,6 @@ class FileHandler(ABC):
139
124
  OSError: If the file contains binary content or has encoding issues.
140
125
  PermissionError: If permission is denied while reading the file.
141
126
  """
142
- logger.debug("Validating file is readable text: %s", str(self.filepath))
143
127
  try:
144
128
  with self.filepath.open("r", encoding=self.ENCODING) as file:
145
129
  # Read first 512 bytes to check for binary content
@@ -147,7 +131,6 @@ class FileHandler(ABC):
147
131
  if "\x00" in sample:
148
132
  logger.error("File contains binary content: %s", str(self.filepath))
149
133
  raise OSError(f"Invalid file format: '{self.filepath}' contains binary data, expected text file")
150
- logger.debug("Text file validation passed: %s", str(self.filepath))
151
134
  except UnicodeDecodeError as e:
152
135
  logger.error("File encoding error (not valid UTF-8): %s - %s", str(self.filepath), e)
153
136
  raise OSError(f"Invalid file encoding: '{self.filepath}' is not valid UTF-8") from e
@@ -169,8 +152,7 @@ class FileHandler(ABC):
169
152
  OSError: If the file fails validation (empty, too large, binary, wrong type).
170
153
  NotImplementedError: If the subclass does not implement `_read()`.
171
154
  """
172
- logger.info("Starting file validation and reading: %s", str(self.filepath))
173
- logger.debug("Running validation checks for file: %s", str(self.filepath))
155
+ logger.info("Reading file: %s", str(self.filepath))
174
156
 
175
157
  self._file_exists()
176
158
  self._is_file()
@@ -179,11 +161,8 @@ class FileHandler(ABC):
179
161
  self._file_size_limits()
180
162
  self._text_file()
181
163
 
182
- logger.info("All validation checks passed for file: %s", str(self.filepath))
183
-
184
- logger.debug("Reading file content: %s", str(self.filepath))
185
164
  data = self._read()
186
- logger.info("File successfully read and parsed: %s", str(self.filepath))
165
+ logger.info("Successfully parsed file: %s", str(self.filepath))
187
166
  return data
188
167
 
189
168
  @abstractmethod
@@ -222,14 +201,9 @@ class FileYamlHandler(FileHandler):
222
201
  FileNotFoundError: If the file does not exist.
223
202
  PermissionError: If the file cannot be read due to permission restrictions.
224
203
  """
225
- logger.info("Reading YAML file: %s", str(self.filepath))
226
-
227
204
  try:
228
- logger.debug("Opening YAML file for reading: %s", str(self.filepath))
229
205
  with open(file=self.filepath, mode="r", encoding="utf-8") as file:
230
206
  data = yaml.safe_load(file)
231
- logger.info("Successfully parsed YAML file: %s", str(self.filepath))
232
- logger.debug("YAML data structure type: %s", type(data))
233
207
  return data
234
208
  except yaml.YAMLError as e:
235
209
  logger.error("YAML parsing error in file '%s': %s", str(self.filepath), e)
@@ -256,15 +230,10 @@ class FileJsonHandler(FileHandler):
256
230
  FileNotFoundError: If the file does not exist.
257
231
  PermissionError: If the file cannot be read due to permission restrictions.
258
232
  """
259
- logger.info("Reading JSON file: %s", str(self.filepath))
260
-
261
233
  try:
262
- logger.debug("Opening JSON file for reading: %s", str(self.filepath))
263
234
  with open(file=self.filepath, mode="r", encoding="utf-8") as file:
264
235
  content = file.read()
265
236
  data = json.loads(content)
266
- logger.info("Successfully parsed JSON file: %s", str(self.filepath))
267
- logger.debug("JSON data structure type: %s", type(data))
268
237
  return data
269
238
  except json.Json5DecoderException as e:
270
239
  logger.error("JSON parsing error in file '%s': %s", str(self.filepath), e)
@@ -316,9 +285,7 @@ class FileHandlerContext:
316
285
  >>> yaml_handler = FileHandlerContext.from_filepath(Path("pipeline.yaml"))
317
286
  >>> json_handler = FileHandlerContext.from_filepath(Path("config.jsonc"))
318
287
  """
319
- logger.debug("Creating file handler for path: %s", str(filepath))
320
288
  _, file_extension = os.path.splitext(filepath)
321
- logger.debug("Detected file extension: %s", file_extension)
322
289
 
323
290
  handler_class = cls.SUPPORTED_EXTENSIONS.get(file_extension)
324
291
 
@@ -335,7 +302,6 @@ class FileHandlerContext:
335
302
  f"Supported formats: {supported_extensions}"
336
303
  )
337
304
 
338
- logger.debug("Selected handler class: %s for extension: %s", handler_class.__name__, file_extension)
339
305
  handler = handler_class(filepath=filepath)
340
306
  logger.info("Created %s for file: %s", handler_class.__name__, str(filepath))
341
307
  return handler
@@ -25,6 +25,7 @@ Typical Usage:
25
25
  import logging
26
26
 
27
27
  import structlog
28
+ from rich.logging import RichHandler
28
29
 
29
30
 
30
31
  def set_logger(level: str = "INFO") -> None:
@@ -77,22 +78,33 @@ def set_logger(level: str = "INFO") -> None:
77
78
  See Also:
78
79
  get_logger: Create logger instances for specific modules
79
80
  """
80
- # Configure standard logging for third-party libraries and OTLP handler
81
- logging.basicConfig(
82
- format="%(message)s",
83
- level=level,
84
- force=True,
81
+ # Configure standard logging for pretty console output
82
+ root_logger = logging.getLogger()
83
+ root_logger.setLevel(level)
84
+
85
+ # Remove existing handlers to avoid duplicates on reconfiguration
86
+ root_logger.handlers.clear()
87
+
88
+ # Use Rich for pretty console output with structlog formatting
89
+ console_handler = RichHandler(rich_tracebacks=True, tracebacks_show_locals=True, markup=True)
90
+ console_handler.setLevel(level)
91
+ console_handler.setFormatter(
92
+ structlog.stdlib.ProcessorFormatter(
93
+ processors=[
94
+ structlog.stdlib.ProcessorFormatter.remove_processors_meta,
95
+ structlog.dev.ConsoleRenderer(colors=False),
96
+ ],
97
+ ),
85
98
  )
99
+ root_logger.addHandler(console_handler)
86
100
 
87
101
  # Configure structlog
88
102
  structlog.configure(
89
103
  processors=[
90
104
  structlog.contextvars.merge_contextvars,
91
105
  structlog.stdlib.add_log_level,
92
- structlog.stdlib.add_logger_name,
93
- structlog.processors.TimeStamper(fmt="iso"),
106
+ structlog.stdlib.PositionalArgumentsFormatter(),
94
107
  structlog.processors.StackInfoRenderer(),
95
- structlog.processors.format_exc_info,
96
108
  structlog.stdlib.ProcessorFormatter.wrap_for_formatter,
97
109
  ],
98
110
  wrapper_class=structlog.stdlib.BoundLogger,
@@ -188,7 +188,6 @@ class WorkflowController(BaseModel):
188
188
  validating external configuration sources, or providing IDE hints
189
189
  for configuration files.
190
190
  """
191
- logger.debug("Exporting WorkflowController JSON schema")
192
191
  return cls.model_json_schema(schema_generator=PreserveFieldOrderJsonSchema)
193
192
 
194
193
  @trace_span("workflow_controller.execute_all")
@@ -8,7 +8,7 @@ how data should be read from various sources with type validation.
8
8
  from enum import Enum
9
9
  from typing import Literal
10
10
 
11
- from pydantic import Field, FilePath
11
+ from pydantic import Field
12
12
 
13
13
  from samara import BaseModel
14
14
  from samara.utils.logger import get_logger
@@ -74,7 +74,7 @@ class ExtractModel(BaseModel):
74
74
  id_: str = Field(..., alias="id", description="Identifier for this extraction operation", min_length=1)
75
75
  method: ExtractMethod = Field(..., description="Method of extraction (batch or streaming)")
76
76
  data_format: str = Field(..., description="Format of the data to extract (parquet, json, csv, etc.)")
77
- schema_: str | FilePath = Field(..., alias="schema", description="Schema definition - can be a file path or string")
77
+ schema_: str = Field(..., alias="schema", description="Schema definition - file path or inline JSON string")
78
78
 
79
79
 
80
80
  class ExtractFileModel(ExtractModel):
@@ -109,9 +109,11 @@ class LoadModel(BaseModel, ABC):
109
109
  upstream_id: str = Field(..., description="Identifier of the upstream component providing data", min_length=1)
110
110
  method: LoadMethod = Field(..., description="Loading method (batch or streaming)")
111
111
  location: str = Field(
112
- ..., description="URI that identifies where to load data in the modelified format.", min_length=1
112
+ ..., description="URI that identifies where to load data in the specified format.", min_length=1
113
+ )
114
+ schema_export: str = Field(
115
+ ..., description="URI where the schema is exported. Use an empty string to skip schema export."
113
116
  )
114
- schema_export: str = Field(..., description="URI that identifies where to load schema.")
115
117
 
116
118
 
117
119
  class LoadModelFile(LoadModel):
@@ -34,7 +34,6 @@ class ArgsModel(BaseModel, ABC):
34
34
 
35
35
 
36
36
  ArgsT = TypeVar("ArgsT", bound=ArgsModel)
37
- FunctionNameT = TypeVar("FunctionNameT", bound=str)
38
37
 
39
38
 
40
39
  class FunctionModel(BaseModel, Generic[ArgsT], ABC):
@@ -7,10 +7,9 @@ schema parsing and PySpark configuration management.
7
7
  """
8
8
 
9
9
  from abc import ABC, abstractmethod
10
- from pathlib import Path
11
10
  from typing import Any, Literal, Self
12
11
 
13
- from pydantic import Field, model_validator
12
+ from pydantic import Field, PrivateAttr, model_validator
14
13
  from pyspark.sql import DataFrame
15
14
  from pyspark.sql.types import StructType
16
15
 
@@ -18,7 +17,7 @@ from samara.telemetry import trace_span
18
17
  from samara.types import DataFrameRegistry
19
18
  from samara.utils.logger import get_logger
20
19
  from samara.workflow.jobs.models.model_extract import ExtractFileModel, ExtractMethod, ExtractModel
21
- from samara.workflow.jobs.spark.schema import SchemaFilepathHandler, SchemaStringHandler
20
+ from samara.workflow.jobs.spark.schema import SchemaHandler
22
21
  from samara.workflow.jobs.spark.session import SparkHandler
23
22
 
24
23
  logger = get_logger(__name__)
@@ -67,27 +66,14 @@ class ExtractSpark(ExtractModel, ABC):
67
66
  extraction mode based on the configured method.
68
67
  """
69
68
 
70
- model_config = {"arbitrary_types_allowed": True, "extra": "allow"}
69
+ model_config = {"arbitrary_types_allowed": True, "extra": "forbid"}
70
+
71
+ _data_registry: DataFrameRegistry = PrivateAttr(default_factory=DataFrameRegistry)
72
+ _spark: SparkHandler = PrivateAttr(default_factory=SparkHandler)
71
73
 
72
74
  _schema_parsed: StructType
73
75
  options: dict[str, Any] = Field(..., description="PySpark reader options as key-value pairs")
74
76
 
75
- def __init__(self, **data: Any) -> None:
76
- """Initialize the extractor with configuration data and workflow components.
77
-
78
- Sets up the Pydantic model with provided configuration data and initializes
79
- workflow components (DataFrameRegistry for storing extracted data and
80
- SparkHandler for managing the Spark session).
81
-
82
- Args:
83
- **data: Configuration data for model initialization. Should include
84
- required fields like id_, method, schema_, and extract-specific fields.
85
- """
86
- super().__init__(**data)
87
- # Set up non-Pydantic attributes that shouldn't be in schema
88
- self.data_registry: DataFrameRegistry = DataFrameRegistry()
89
- self.spark: SparkHandler = SparkHandler()
90
-
91
77
  @model_validator(mode="after")
92
78
  @trace_span("extract_spark.parse_schema")
93
79
  def parse_schema(self) -> Self:
@@ -103,21 +89,13 @@ class ExtractSpark(ExtractModel, ABC):
103
89
 
104
90
  Note:
105
91
  If schema_ is empty or None, returns early without parsing.
106
- File path detection relies on the .json extension.
92
+ Detection uses JSON parsing: valid JSON is treated as inline,
93
+ anything else is treated as a file path.
107
94
  """
108
95
  if not self.schema_:
109
96
  return self
110
97
 
111
- # Convert to string for processing
112
- schema_str = str(self.schema_).strip()
113
-
114
- # Detect if it's a file path or JSON string
115
- if schema_str.endswith(".json"):
116
- # File path - use FilepathHandler
117
- self._schema_parsed = SchemaFilepathHandler.parse(schema=Path(schema_str))
118
- else:
119
- # JSON string - use StringHandler
120
- self._schema_parsed = SchemaStringHandler.parse(schema=schema_str)
98
+ self._schema_parsed = SchemaHandler.parse_schema_string(schema=str(self.schema_).strip())
121
99
 
122
100
  return self
123
101
 
@@ -139,19 +117,15 @@ class ExtractSpark(ExtractModel, ABC):
139
117
  """
140
118
  logger.info("Starting extraction for source: %s using method: %s", self.id_, self.method.value)
141
119
 
142
- logger.debug("Adding Spark configurations: %s", self.options)
143
- self.spark.add_configs(options=self.options)
144
-
145
- if self.method == ExtractMethod.BATCH:
146
- logger.debug("Performing batch extraction for: %s", self.id_)
147
- self.data_registry[self.id_] = self._extract_batch()
148
- logger.info("Batch extraction completed successfully for: %s", self.id_)
149
- elif self.method == ExtractMethod.STREAMING:
150
- logger.debug("Performing streaming extraction for: %s", self.id_)
151
- self.data_registry[self.id_] = self._extract_streaming()
152
- logger.info("Streaming extraction completed successfully for: %s", self.id_)
153
- else:
154
- raise ValueError(f"Extraction method {self.method} is not supported for PySpark")
120
+ with self._spark.scoped_configs(options=self.options):
121
+ if self.method == ExtractMethod.BATCH:
122
+ self._data_registry[self.id_] = self._extract_batch()
123
+ logger.info("Batch extraction completed for: %s", self.id_)
124
+ elif self.method == ExtractMethod.STREAMING:
125
+ self._data_registry[self.id_] = self._extract_streaming()
126
+ logger.info("Streaming extraction completed for: %s", self.id_)
127
+ else:
128
+ raise ValueError(f"Extraction method {self.method} is not supported for PySpark")
155
129
 
156
130
  @abstractmethod
157
131
  def _extract_batch(self) -> DataFrame:
@@ -245,14 +219,13 @@ class ExtractFileSpark(ExtractSpark, ExtractFileModel):
245
219
  """
246
220
  logger.debug("Reading files in batch mode - path: %s, format: %s", self.location, self.data_format)
247
221
 
248
- dataframe = self.spark.session.read.load(
222
+ dataframe = self._spark.session.read.load(
249
223
  path=self.location,
250
224
  format=self.data_format,
251
225
  schema=self._schema_parsed,
252
226
  **self.options,
253
227
  )
254
- row_count = dataframe.count()
255
- logger.info("Batch extraction successful - loaded %d rows from %s", row_count, self.location)
228
+ logger.info("Batch extraction successful - loaded data from %s", self.location)
256
229
  return dataframe
257
230
 
258
231
  @trace_span("extract_file_spark._extract_streaming")
@@ -272,7 +245,7 @@ class ExtractFileSpark(ExtractSpark, ExtractFileModel):
272
245
  """
273
246
  logger.debug("Reading files in streaming mode - path: %s, format: %s", self.location, self.data_format)
274
247
 
275
- dataframe = self.spark.session.readStream.load(
248
+ dataframe = self._spark.session.readStream.load(
276
249
  path=self.location,
277
250
  format=self.data_format,
278
251
  schema=self._schema_parsed,
@@ -16,6 +16,7 @@ from samara.utils.logger import get_logger
16
16
  from samara.workflow.jobs.models.model_job import JobEngine, JobModel
17
17
  from samara.workflow.jobs.spark.extract import ExtractSparkUnion
18
18
  from samara.workflow.jobs.spark.load import LoadSparkUnion
19
+ from samara.workflow.jobs.spark.session import SparkHandler
19
20
  from samara.workflow.jobs.spark.transform import TransformSparkUnion
20
21
 
21
22
  logger = get_logger(__name__)
@@ -159,7 +160,7 @@ class JobSpark(JobModel[ExtractSparkUnion, TransformSparkUnion, LoadSparkUnion])
159
160
  logger.info("Starting extract phase with %d extractors", len(self.extracts))
160
161
  start_time = time.time()
161
162
 
162
- for i, extract in enumerate(self.extracts):
163
+ for i, extract in enumerate(self.extracts, 1):
163
164
  extract_start_time = time.time()
164
165
  logger.debug("Running extractor %d/%d: %s", i, len(self.extracts), extract.id_)
165
166
  extract.extract()
@@ -189,7 +190,7 @@ class JobSpark(JobModel[ExtractSparkUnion, TransformSparkUnion, LoadSparkUnion])
189
190
  logger.info("Starting transform phase with %d transformers", len(self.transforms))
190
191
  start_time = time.time()
191
192
 
192
- for i, transform in enumerate(self.transforms):
193
+ for i, transform in enumerate(self.transforms, 1):
193
194
  transform_start_time = time.time()
194
195
  logger.debug("Running transformer %d/%d: %s", i, len(self.transforms), transform.id_)
195
196
  transform.transform()
@@ -220,7 +221,7 @@ class JobSpark(JobModel[ExtractSparkUnion, TransformSparkUnion, LoadSparkUnion])
220
221
  logger.info("Starting load phase with %d loaders", len(self.loads))
221
222
  start_time = time.time()
222
223
 
223
- for i, load in enumerate(self.loads):
224
+ for i, load in enumerate(self.loads, 1):
224
225
  load_start_time = time.time()
225
226
  logger.debug("Running loader %d/%d: %s", i, len(self.loads), load.id_)
226
227
  load.load()
@@ -232,15 +233,13 @@ class JobSpark(JobModel[ExtractSparkUnion, TransformSparkUnion, LoadSparkUnion])
232
233
 
233
234
  @override
234
235
  def _clear(self) -> None:
235
- """Free resources by clearing Spark-specific registries.
236
+ """Free resources by clearing Spark-specific registries and stopping the session.
236
237
 
237
238
  Clears the DataFrameRegistry and StreamingQueryRegistry after job execution
238
- completes. This prevents memory leaks and ensures clean state for subsequent
239
- jobs, particularly important in long-running processes or batch environments
240
- where multiple jobs execute sequentially.
239
+ completes, then stops the SparkSession to release JVM resources. This prevents
240
+ memory leaks and ensures clean state for subsequent jobs, particularly important
241
+ in long-running processes or containerized environments.
241
242
  """
242
- logger.debug("Clearing DataFrameRegistry after job: %s", self.id_)
243
243
  DataFrameRegistry().clear()
244
-
245
- logger.debug("Clearing StreamingQueryRegistry after job: %s", self.id_)
246
244
  StreamingQueryRegistry().clear()
245
+ SparkHandler().stop_session()
@@ -9,7 +9,7 @@ import json
9
9
  from abc import ABC, abstractmethod
10
10
  from typing import Any, Literal
11
11
 
12
- from pydantic import Field
12
+ from pydantic import Field, PrivateAttr
13
13
  from pyspark.sql.streaming.query import StreamingQuery
14
14
 
15
15
  from samara.telemetry import trace_span
@@ -89,29 +89,13 @@ class LoadSpark(LoadModel, ABC):
89
89
  Spark application.
90
90
  """
91
91
 
92
- model_config = {"arbitrary_types_allowed": True, "extra": "allow"}
92
+ model_config = {"arbitrary_types_allowed": True, "extra": "forbid"}
93
93
 
94
- options: dict[str, Any] = Field(..., description="Options for the sink input.")
95
-
96
- def __init__(self, **data: Any) -> None:
97
- """Initialize the load component from pipeline configuration.
98
-
99
- Creates a Pydantic model instance from provided configuration and
100
- initializes workflow registries and Spark handler. Establishes the
101
- component's ability to manage DataFrames, streaming queries, and
102
- Spark session settings during the load operation.
94
+ _data_registry: DataFrameRegistry = PrivateAttr(default_factory=DataFrameRegistry)
95
+ _streaming_query_registry: StreamingQueryRegistry = PrivateAttr(default_factory=StreamingQueryRegistry)
96
+ _spark: SparkHandler = PrivateAttr(default_factory=SparkHandler)
103
97
 
104
- Args:
105
- **data: Configuration fields from pipeline definition. Expected
106
- fields include: id_, upstream_id, method, location,
107
- data_format, mode, options, and optionally schema_export.
108
- See class Example for full configuration structure.
109
- """
110
- super().__init__(**data)
111
- # Set up non-Pydantic attributes that shouldn't be in schema
112
- self.data_registry: DataFrameRegistry = DataFrameRegistry()
113
- self.streaming_query_registry: StreamingQueryRegistry = StreamingQueryRegistry()
114
- self.spark: SparkHandler = SparkHandler()
98
+ options: dict[str, Any] = Field(..., description="Options for the sink input.")
115
99
 
116
100
  @abstractmethod
117
101
  def _load_batch(self) -> None:
@@ -181,29 +165,21 @@ class LoadSpark(LoadModel, ABC):
181
165
  self.method.value,
182
166
  )
183
167
 
184
- logger.debug("Adding Spark configurations: %s", self.options)
185
- self.spark.add_configs(options=self.options)
186
-
187
- logger.debug("Copying dataframe from %s to %s", self.upstream_id, self.id_)
188
- self.data_registry[self.id_] = self.data_registry[self.upstream_id]
168
+ with self._spark.scoped_configs(options=self.options):
169
+ self._data_registry[self.id_] = self._data_registry[self.upstream_id]
189
170
 
190
- if self.method == LoadMethod.BATCH:
191
- logger.debug("Performing batch load for: %s", self.id_)
192
- self._load_batch()
193
- logger.info("Batch load completed successfully for: %s", self.id_)
194
- elif self.method == LoadMethod.STREAMING:
195
- logger.debug("Performing streaming load for: %s", self.id_)
196
- self.streaming_query_registry[self.id_] = self._load_streaming()
197
- logger.info("Streaming load started successfully for: %s", self.id_)
198
- else:
199
- raise ValueError(f"Loading method {self.method} is not supported for PySpark")
171
+ if self.method == LoadMethod.BATCH:
172
+ self._load_batch()
173
+ logger.info("Batch load completed for: %s", self.id_)
174
+ elif self.method == LoadMethod.STREAMING:
175
+ self._streaming_query_registry[self.id_] = self._load_streaming()
176
+ logger.info("Streaming load started for: %s", self.id_)
177
+ else:
178
+ raise ValueError(f"Loading method {self.method} is not supported for PySpark")
200
179
 
201
- # Export schema if location is specified
202
- if self.schema_export:
203
- schema_json = json.dumps(self.data_registry[self.id_].schema.jsonValue())
204
- self._export_schema(schema_json, self.schema_export)
205
-
206
- logger.info("Load operation completed successfully for: %s", self.id_)
180
+ if self.schema_export:
181
+ schema_json = json.dumps(self._data_registry[self.id_].schema.jsonValue())
182
+ self._export_schema(schema_json, self.schema_export)
207
183
 
208
184
 
209
185
  class LoadFileSpark(LoadSpark, LoadModelFile):
@@ -289,17 +265,14 @@ class LoadFileSpark(LoadSpark, LoadModelFile):
289
265
  self.mode,
290
266
  )
291
267
 
292
- row_count = self.data_registry[self.id_].count()
293
- logger.debug("Writing %d rows to %s", row_count, self.location)
294
-
295
- self.data_registry[self.id_].write.save(
268
+ self._data_registry[self.id_].write.save(
296
269
  path=self.location,
297
270
  format=self.data_format,
298
271
  mode=self.mode,
299
272
  **self.options,
300
273
  )
301
274
 
302
- logger.info("Batch write successful - wrote %d rows to %s", row_count, self.location)
275
+ logger.info("Batch write successful - wrote data to %s", self.location)
303
276
 
304
277
  @trace_span("load_file_spark._load_streaming")
305
278
  def _load_streaming(self) -> StreamingQuery:
@@ -320,7 +293,7 @@ class LoadFileSpark(LoadSpark, LoadModelFile):
320
293
  self.mode,
321
294
  )
322
295
 
323
- streaming_query = self.data_registry[self.id_].writeStream.start(
296
+ streaming_query = self._data_registry[self.id_].writeStream.start(
324
297
  path=self.location,
325
298
  format=self.data_format,
326
299
  outputMode=self.mode,
@@ -53,6 +53,24 @@ class SchemaHandler(ABC):
53
53
  NotImplementedError: If not implemented by a subclass.
54
54
  """
55
55
 
56
+ @staticmethod
57
+ def parse_schema_string(schema: str) -> StructType:
58
+ """Parse a schema string, auto-detecting whether it is inline JSON or a file path.
59
+
60
+ A PySpark schema definition is always a JSON object starting with '{'.
61
+ Anything else is treated as a file path.
62
+
63
+ Args:
64
+ schema: A schema string that is either a JSON object defining an
65
+ inline schema or a file path pointing to a schema definition file.
66
+
67
+ Returns:
68
+ StructType: A fully configured PySpark StructType schema.
69
+ """
70
+ if schema.startswith("{"):
71
+ return SchemaStringHandler.parse(schema=schema)
72
+ return SchemaFilepathHandler.parse(schema=Path(schema))
73
+
56
74
 
57
75
  class SchemaDictHandler(SchemaHandler):
58
76
  """Convert dictionary schemas to PySpark StructType.
@@ -101,9 +119,7 @@ class SchemaDictHandler(SchemaHandler):
101
119
 
102
120
  try:
103
121
  struct_type = StructType.fromJson(json=schema)
104
- field_count = len(struct_type.fields)
105
- logger.info("Successfully parsed schema from dictionary - %d fields", field_count)
106
- logger.debug("Schema fields: %s", [f.name for f in struct_type.fields])
122
+ logger.info("Successfully parsed schema from dictionary - %d fields", len(struct_type.fields))
107
123
  return struct_type
108
124
  except (ValueError, TypeError, KeyError) as e:
109
125
  raise ValueError(f"Failed to convert dictionary to schema: {e}") from e
@@ -151,15 +167,9 @@ class SchemaStringHandler(SchemaHandler):
151
167
  valid schema structure.
152
168
  json.JSONDecodeError: If the string is not valid JSON.
153
169
  """
154
- logger.debug("Parsing schema from JSON string (length: %d)", len(schema))
155
-
156
170
  try:
157
- logger.debug("Parsing JSON string to dictionary")
158
171
  parsed_json = json.loads(s=schema)
159
- logger.debug("Successfully parsed JSON string")
160
-
161
172
  result = SchemaDictHandler.parse(schema=parsed_json)
162
- logger.info("Successfully parsed schema from JSON string")
163
173
  return result
164
174
 
165
175
  except json.JSONDecodeError as e:
@@ -211,13 +221,8 @@ class SchemaFilepathHandler(SchemaHandler):
211
221
  logger.info("Parsing schema from file: %s", str(schema))
212
222
 
213
223
  try:
214
- logger.debug("Creating file handler for schema file: %s", str(schema))
215
224
  file_handler: FileHandler = FileHandlerContext.from_filepath(filepath=schema)
216
-
217
- logger.debug("Reading schema file content")
218
225
  file_content = file_handler.read()
219
-
220
- logger.debug("Converting file content to schema")
221
226
  result = SchemaDictHandler.parse(schema=file_content)
222
227
 
223
228
  logger.info("Successfully parsed schema from file: %s", str(schema))
@@ -8,12 +8,15 @@ Key features:
8
8
  - Lazy initialization of Spark sessions (only created when needed)
9
9
  - Automatic resource cleanup and management
10
10
  - Centralized configuration handling for Spark parameters
11
+ - Stage-scoped configuration to prevent config leaking between pipeline stages
11
12
  - Seamless integration with Samara's configuration-driven pipeline model
12
13
 
13
14
  The SparkHandler singleton ensures efficient resource usage across the entire
14
15
  pipeline lifecycle, whether running locally for testing or on distributed clusters.
15
16
  """
16
17
 
18
+ from collections.abc import Iterator
19
+ from contextlib import contextmanager
17
20
  from typing import Any
18
21
 
19
22
  from pyspark.sql import SparkSession
@@ -67,7 +70,6 @@ class SparkHandler(metaclass=Singleton):
67
70
  options: Spark configuration from your pipeline definition as key-value
68
71
  pairs (e.g., {"spark.executor.memory": "4g"}). Optional.
69
72
  """
70
- logger.debug("Configuring SparkHandler with app_name: %s (lazy initialization)", app_name)
71
73
  self._session = None
72
74
  self._app_name = app_name
73
75
  self._init_options = options or {}
@@ -84,7 +86,7 @@ class SparkHandler(metaclass=Singleton):
84
86
  The Spark session ready to execute your transformations
85
87
  """
86
88
  if self._session is None:
87
- logger.debug("Creating SparkSession on first access - app_name: %s", self._app_name)
89
+ logger.debug("Creating SparkSession - app_name: %s", self._app_name)
88
90
 
89
91
  builder = SparkSession.Builder().appName(name=self._app_name)
90
92
 
@@ -93,11 +95,9 @@ class SparkHandler(metaclass=Singleton):
93
95
  logger.debug("Setting Spark config: %s = %s", key, value)
94
96
  builder = builder.config(key=key, value=value)
95
97
 
96
- logger.debug("Creating/retrieving SparkSession")
97
98
  self._session = builder.getOrCreate()
98
- logger.info("SparkHandler initialized successfully with app: %s", self._app_name)
99
+ logger.info("SparkSession initialized with app: %s", self._app_name)
99
100
 
100
- logger.debug("Accessing SparkSession instance")
101
101
  return self._session
102
102
 
103
103
  @session.setter
@@ -126,6 +126,14 @@ class SparkHandler(metaclass=Singleton):
126
126
  Use this to manually clean up if you need to restart Spark during a
127
127
  pipeline's lifecycle.
128
128
  """
129
+ self.stop_session()
130
+
131
+ def stop_session(self) -> None:
132
+ """Stop the active Spark session and release all associated resources.
133
+
134
+ Safely shuts down the Spark session if one is active. Idempotent:
135
+ calling on an already-stopped or never-started session is a no-op.
136
+ """
129
137
  if self._session is not None:
130
138
  logger.info("Stopping SparkSession: %s", self._session.sparkContext.appName)
131
139
  self._session.stop()
@@ -149,6 +157,10 @@ class SparkHandler(metaclass=Singleton):
149
157
  Some Spark settings cannot be changed after initialization. For
150
158
  pre-execution configuration, define settings in your pipeline's
151
159
  engine configuration instead.
160
+
161
+ Configs applied here persist for the lifetime of the SparkSession.
162
+ For stage-scoped configs that automatically revert, use
163
+ scoped_configs() instead.
152
164
  """
153
165
  logger.debug("Adding %d configuration options to SparkSession", len(options))
154
166
 
@@ -157,3 +169,40 @@ class SparkHandler(metaclass=Singleton):
157
169
  self.session.conf.set(key=key, value=value)
158
170
 
159
171
  logger.info("Successfully applied %d configuration options", len(options))
172
+
173
+ @contextmanager
174
+ def scoped_configs(self, options: dict[str, Any]) -> Iterator[None]:
175
+ """Apply Spark settings for the duration of a pipeline stage, then restore previous values.
176
+
177
+ Context manager that temporarily applies Spark configuration options and
178
+ automatically reverts them when the stage completes. This prevents configs
179
+ set by one stage from leaking into subsequent stages.
180
+
181
+ Args:
182
+ options: Configuration settings as key-value pairs to apply
183
+ temporarily (e.g., {"spark.sql.shuffle.partitions": "200"})
184
+
185
+ Yields:
186
+ None: Control returns to the caller with configs applied.
187
+ """
188
+ if not options:
189
+ yield
190
+ return
191
+
192
+ logger.debug("Applying %d scoped configuration options", len(options))
193
+
194
+ previous_values: dict[str, str | None] = {}
195
+ for key in options:
196
+ previous_values[key] = self.session.conf.get(key, default=None)
197
+
198
+ self.add_configs(options)
199
+
200
+ try:
201
+ yield
202
+ finally:
203
+ logger.debug("Restoring %d configuration options to previous values", len(previous_values))
204
+ for key, old_value in previous_values.items():
205
+ if old_value is None:
206
+ self.session.conf.unset(key)
207
+ else:
208
+ self.session.conf.set(key=key, value=old_value)
@@ -8,7 +8,7 @@ through structured configuration rather than code.
8
8
 
9
9
  from typing import Any
10
10
 
11
- from pydantic import Field
11
+ from pydantic import Field, PrivateAttr
12
12
 
13
13
  from samara.telemetry import trace_span
14
14
  from samara.types import DataFrameRegistry
@@ -89,28 +89,12 @@ class TransformSpark(TransformModel[TransformFunctionSparkUnion]):
89
89
  function modifies the dataframe in place within the registry, so order matters.
90
90
  """
91
91
 
92
- model_config = {"arbitrary_types_allowed": True, "extra": "allow"}
92
+ model_config = {"arbitrary_types_allowed": True, "extra": "forbid"}
93
93
 
94
- options: dict[str, Any] = Field(..., description="Transformation options as key-value pairs")
95
-
96
- def __init__(self, **data: Any) -> None:
97
- """Initialize TransformSpark with configuration data.
98
-
99
- Creates a Pydantic model instance from the provided configuration data and then
100
- initializes workflow attributes for managing dataframes and Spark sessions. This
101
- two-stage initialization separates Pydantic model validation from workflow setup.
102
-
103
- Args:
104
- **data: Configuration data for initializing the Pydantic model. Should include
105
- `id_`, `upstream_id`, `functions`, and `options` keys at minimum.
94
+ _data_registry: DataFrameRegistry = PrivateAttr(default_factory=DataFrameRegistry)
95
+ _spark: SparkHandler = PrivateAttr(default_factory=SparkHandler)
106
96
 
107
- Returns:
108
- None
109
- """
110
- super().__init__(**data)
111
- # Set up non-Pydantic attributes that shouldn't be in schema
112
- self.data_registry: DataFrameRegistry = DataFrameRegistry()
113
- self.spark: SparkHandler = SparkHandler()
97
+ options: dict[str, Any] = Field(..., description="Transformation options as key-value pairs")
114
98
 
115
99
  @trace_span("transform_spark.transform")
116
100
  def transform(self) -> None:
@@ -133,29 +117,18 @@ class TransformSpark(TransformModel[TransformFunctionSparkUnion]):
133
117
  """
134
118
  logger.info("Starting transformation for: %s from upstream: %s", self.id_, self.upstream_id)
135
119
 
136
- logger.debug("Adding Spark configurations: %s", self.options)
137
- self.spark.add_configs(options=self.options)
138
-
139
- # Copy the dataframe from upstream to current id
140
- logger.debug("Copying dataframe from %s to %s", self.upstream_id, self.id_)
141
- self.data_registry[self.id_] = self.data_registry[self.upstream_id]
142
-
143
- # Apply transformations
144
- logger.debug("Applying %d transformation functions", len(self.functions))
145
- for i, function in enumerate(self.functions):
146
- logger.debug("Applying function %d/%d: %s", i, len(self.functions), function.function_type)
120
+ with self._spark.scoped_configs(options=self.options):
121
+ self._data_registry[self.id_] = self._data_registry[self.upstream_id]
147
122
 
148
- original_count = self.data_registry[self.id_].count()
149
- callable_ = function.transform()
150
- self.data_registry[self.id_] = callable_(df=self.data_registry[self.id_])
123
+ for i, function in enumerate(self.functions, 1):
124
+ logger.debug("Applying function %d/%d: %s", i, len(self.functions), function.function_type)
151
125
 
152
- new_count = self.data_registry[self.id_].count()
126
+ callable_ = function.transform()
127
+ self._data_registry[self.id_] = callable_(df=self._data_registry[self.id_])
153
128
 
154
- logger.info(
155
- "Function %s applied - rows changed from %d to %d", function.function_type, original_count, new_count
156
- )
129
+ logger.info("Function %s applied successfully", function.function_type)
157
130
 
158
- logger.info("Transformation completed successfully for: %s", self.id_)
131
+ logger.info("Transformation completed successfully for: %s", self.id_)
159
132
 
160
133
 
161
134
  TransformSparkUnion = TransformSpark
@@ -4,7 +4,7 @@ This module provides the base class for all Spark-specific transformation functi
4
4
  enabling shared access to the DataFrame registry across transformation operations.
5
5
  """
6
6
 
7
- from typing import ClassVar
7
+ from pydantic import PrivateAttr
8
8
 
9
9
  from samara.types import DataFrameRegistry
10
10
  from samara.workflow.jobs.models.model_transform import ArgsT, FunctionModel
@@ -20,14 +20,14 @@ class FunctionSpark(FunctionModel[ArgsT]):
20
20
  registry access throughout the transformation execution.
21
21
 
22
22
  Attributes:
23
- data_registry: Shared class-level registry for accessing processed
24
- DataFrames by their identifier within the pipeline execution context.
23
+ _data_registry: Shared registry for accessing processed DataFrames by
24
+ their identifier within the pipeline execution context.
25
25
 
26
26
  Note:
27
- The data_registry is a class-level attribute shared across all instances
28
- within a pipeline execution, enabling cross-reference between DataFrames
29
- created by different transformation steps. This is essential for operations
30
- that operate on multiple DataFrames such as joins and unions.
27
+ The _data_registry is a private attribute initialized per-instance via
28
+ Pydantic's PrivateAttr. Since DataFrameRegistry is a singleton, all
29
+ instances share the same underlying registry, enabling cross-reference
30
+ between DataFrames created by different transformation steps.
31
31
  """
32
32
 
33
- data_registry: ClassVar[DataFrameRegistry] = DataFrameRegistry()
33
+ _data_registry: DataFrameRegistry = PrivateAttr(default_factory=DataFrameRegistry)
@@ -121,15 +121,9 @@ class FilterFunction(FilterFunctionModel, FunctionSpark):
121
121
  """
122
122
 
123
123
  def __f(df: DataFrame) -> DataFrame:
124
- logger.debug("Applying filter transform with condition: %s", self.arguments.condition)
125
- original_count = df.count()
126
- logger.debug("Input DataFrame has %d rows", original_count)
127
-
128
124
  result_df = df.filter(self.arguments.condition)
129
- filtered_count = result_df.count()
130
- filtered_out = original_count - filtered_count
131
125
 
132
- logger.info("Filter transform completed - kept %d rows, filtered out %d rows", filtered_count, filtered_out)
126
+ logger.info("Filter transform completed - condition: %s", self.arguments.condition)
133
127
  return result_df
134
128
 
135
129
  return __f
@@ -90,36 +90,17 @@ class JoinFunction(JoinFunctionModel, FunctionSpark):
90
90
  may require appropriate cluster resources and shuffle operations.
91
91
  """
92
92
  logger.debug(
93
- "Creating join transform - other: %s, on: %s, how: %s",
93
+ "Configuring join - other: %s, on: %s, how: %s",
94
94
  self.arguments.other_upstream_id,
95
95
  self.arguments.on,
96
96
  self.arguments.how,
97
97
  )
98
98
 
99
99
  def __f(df: DataFrame) -> DataFrame:
100
- logger.debug("Applying join transform")
100
+ right_df = self._data_registry[self.arguments.other_upstream_id]
101
+ result_df = df.join(right_df, on=self.arguments.on, how=self.arguments.how)
101
102
 
102
- # Get the right DataFrame from the registry
103
- right_df = self.data_registry[self.arguments.other_upstream_id]
104
- logger.debug(
105
- "Retrieved right DataFrame: %s (columns: %s)",
106
- self.arguments.other_upstream_id,
107
- right_df.columns,
108
- )
109
-
110
- # Get the join type
111
- join_type = self.arguments.how
112
- # Get the join columns
113
- join_on = self.arguments.on
114
-
115
- logger.debug("Performing join - left: %d rows, right: %d rows", df.count(), right_df.count())
116
- logger.debug("Join parameters - on: %s, how: %s", join_on, join_type)
117
-
118
- # Perform the join operation
119
- result_df = df.join(right_df, on=join_on, how=join_type)
120
- result_count = result_df.count()
121
-
122
- logger.info("Join transform completed - result: %d rows, join type: %s", result_count, join_type)
103
+ logger.info("Join transform completed - type: %s", self.arguments.how)
123
104
 
124
105
  return result_df
125
106
 
@@ -102,17 +102,12 @@ class SelectFunction(SelectFunctionModel, FunctionSpark):
102
102
 
103
103
  The output contains only the projected columns in the order specified.
104
104
  """
105
- logger.debug("Creating select transform for columns: %s", self.arguments.columns)
106
105
 
107
106
  def __f(df: DataFrame) -> DataFrame:
108
- logger.debug("Applying select transform - input columns: %s", df.columns)
109
- logger.debug("Selecting columns: %s", self.arguments.columns)
110
-
111
107
  result_df = df.select(*self.arguments.columns)
112
108
  logger.info(
113
109
  "Select transform completed - selected %d columns from %d", len(result_df.columns), len(df.columns)
114
110
  )
115
- logger.debug("Selected columns: %s", result_df.columns)
116
111
  return result_df
117
112
 
118
113
  return __f
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes