streamcase 0.1.0a1__py3-none-any.whl

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.
streamcase/__init__.py ADDED
@@ -0,0 +1,23 @@
1
+ """Deterministic testing for Apache Spark Structured Streaming."""
2
+
3
+ from streamcase._version import __version__
4
+ from streamcase.actions import Batch, Restart, batch, restart
5
+ from streamcase.assertions import assert_batch_count, assert_rows_equal, assert_unique_keys
6
+ from streamcase.results import CapturedBatch, ScenarioResult
7
+ from streamcase.scenario import Action, Scenario, scenario
8
+
9
+ __all__ = [
10
+ "Action",
11
+ "Batch",
12
+ "CapturedBatch",
13
+ "Restart",
14
+ "Scenario",
15
+ "ScenarioResult",
16
+ "__version__",
17
+ "assert_batch_count",
18
+ "assert_rows_equal",
19
+ "assert_unique_keys",
20
+ "batch",
21
+ "restart",
22
+ "scenario",
23
+ ]
streamcase/_cleanup.py ADDED
@@ -0,0 +1,33 @@
1
+ """Preserve execution failures while attempting runner-owned cleanup."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Callable, Iterator
6
+ from contextlib import contextmanager
7
+
8
+
9
+ def _record_cleanup_failure(primary: BaseException, resource: str, error: BaseException) -> None:
10
+ failures = primary.__dict__.get("_streamcase_cleanup_failures", ())
11
+ primary.__dict__["_streamcase_cleanup_failures"] = (*failures, (resource, error))
12
+
13
+ note = f"Streamcase cleanup failed for {resource}: {type(error).__name__}: {error}"
14
+ add_note = getattr(primary, "add_note", None)
15
+ if callable(add_note):
16
+ add_note(note)
17
+ else:
18
+ primary.args = (*primary.args, note)
19
+
20
+
21
+ @contextmanager
22
+ def _cleanup_on_exit(resource: str, cleanup: Callable[[], None]) -> Iterator[None]:
23
+ """Run cleanup and keep a primary failure when cleanup also fails."""
24
+ try:
25
+ yield
26
+ except BaseException as primary:
27
+ try:
28
+ cleanup()
29
+ except BaseException as error:
30
+ _record_cleanup_failure(primary, resource, error)
31
+ raise
32
+ else:
33
+ cleanup()
@@ -0,0 +1,108 @@
1
+ """Deterministic internal formatting for assertion failures."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Mapping, Sequence
6
+ from typing import TypeAlias, cast
7
+
8
+ _MAX_DIAGNOSTIC_ITEMS = 5
9
+
10
+ _Location: TypeAlias = tuple[int, int]
11
+ _DuplicateKey: TypeAlias = tuple[
12
+ tuple[str, ...],
13
+ tuple[object, ...],
14
+ Sequence[_Location],
15
+ ]
16
+
17
+
18
+ def _render_value(value: object) -> str:
19
+ if isinstance(value, Mapping):
20
+ mapping = cast(Mapping[str, object], value)
21
+ rendered_mapping_items = (
22
+ f"{key!r}: {_render_value(mapping[key])}" for key in sorted(mapping)
23
+ )
24
+ return "{" + ", ".join(rendered_mapping_items) + "}"
25
+
26
+ if isinstance(value, tuple):
27
+ rendered_tuple_items = ", ".join(_render_value(item) for item in value)
28
+ trailing_comma = "," if len(value) == 1 else ""
29
+ return f"({rendered_tuple_items}{trailing_comma})"
30
+
31
+ return repr(value)
32
+
33
+
34
+ def _format_item_section(title: str, items: Sequence[str]) -> list[str]:
35
+ lines = [f"{title} ({len(items)}):"]
36
+ visible_items = items[:_MAX_DIAGNOSTIC_ITEMS]
37
+ lines.extend(f" - {item}" for item in visible_items)
38
+
39
+ omitted_count = len(items) - len(visible_items)
40
+ if omitted_count:
41
+ lines.append(f" ... {omitted_count} more item(s) omitted")
42
+ return lines
43
+
44
+
45
+ def _format_batch_count_mismatch(expected_count: int, actual_count: int) -> str:
46
+ return "\n".join(
47
+ (
48
+ "Captured batch count differs.",
49
+ f"Expected: {expected_count} batch(es).",
50
+ f"Actual: {actual_count} batch(es).",
51
+ ),
52
+ )
53
+
54
+
55
+ def _format_rows_mismatch(
56
+ expected_count: int,
57
+ actual_count: int,
58
+ missing: Sequence[Mapping[str, object]],
59
+ unexpected: Sequence[Mapping[str, object]],
60
+ ) -> str:
61
+ lines = [
62
+ "Rows differ.",
63
+ f"Expected: {expected_count} row(s).",
64
+ f"Actual: {actual_count} row(s).",
65
+ ]
66
+ if missing:
67
+ rendered_missing = [_render_value(row) for row in missing]
68
+ lines.extend(_format_item_section("Missing rows", rendered_missing))
69
+ if unexpected:
70
+ rendered_unexpected = [_render_value(row) for row in unexpected]
71
+ lines.extend(_format_item_section("Unexpected rows", rendered_unexpected))
72
+ return "\n".join(lines)
73
+
74
+
75
+ def _format_missing_key_fields(
76
+ batch_id: int,
77
+ row_index: int,
78
+ missing_fields: Sequence[str],
79
+ ) -> str:
80
+ items = [f"batch {batch_id} row {row_index}: {field!r}" for field in missing_fields]
81
+ lines = ["Unique-key assertion failed."]
82
+ lines.extend(_format_item_section("Missing key fields", items))
83
+ return "\n".join(lines)
84
+
85
+
86
+ def _format_duplicate_keys(duplicates: Sequence[_DuplicateKey]) -> str:
87
+ items: list[str] = []
88
+ for fields, key, locations in duplicates:
89
+ rendered_key = ", ".join(
90
+ f"{field}={_render_value(value)}" for field, value in zip(fields, key, strict=True)
91
+ )
92
+
93
+ visible_locations = locations[:_MAX_DIAGNOSTIC_ITEMS]
94
+ rendered_locations = ", ".join(
95
+ f"batch {batch_id} row {row_index}" for batch_id, row_index in visible_locations
96
+ )
97
+ omitted_count = len(locations) - len(visible_locations)
98
+ if omitted_count:
99
+ rendered_locations += f", ... {omitted_count} more occurrence(s) omitted"
100
+
101
+ items.append(
102
+ f"Duplicate key ({rendered_key}) occurred {len(locations)} times at "
103
+ f"{rendered_locations}."
104
+ )
105
+
106
+ lines = ["Unique-key assertion failed."]
107
+ lines.extend(_format_item_section("Duplicate keys", items))
108
+ return "\n".join(lines)
@@ -0,0 +1,107 @@
1
+ """Private isolated filesystem layouts for scenario runs."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ import shutil
7
+ import tempfile
8
+ from dataclasses import dataclass, field
9
+ from pathlib import Path
10
+
11
+ _RUN_DIRECTORY_PREFIX = "streamcase-run-"
12
+
13
+
14
+ def _resolve_base_directory(base_dir: str | os.PathLike[str]) -> Path:
15
+ try:
16
+ raw_path = os.fspath(base_dir)
17
+ except TypeError as error:
18
+ raise TypeError("Runner base directory must be a string or path-like object.") from error
19
+
20
+ if not isinstance(raw_path, str):
21
+ raise TypeError("Runner base directory must resolve to a string path.")
22
+ if not raw_path:
23
+ raise ValueError("Runner base directory must not be empty.")
24
+
25
+ path = Path(raw_path)
26
+ if not path.exists():
27
+ raise FileNotFoundError(f"Runner base directory does not exist: {path}")
28
+ if not path.is_dir():
29
+ raise NotADirectoryError(f"Runner base path is not a directory: {path}")
30
+ return path.resolve(strict=True)
31
+
32
+
33
+ def _require_direct_child(path: Path, parent: Path, description: str) -> None:
34
+ if path.parent != parent:
35
+ raise RuntimeError(f"Generated {description} must be a direct child of {parent}: {path}")
36
+
37
+
38
+ @dataclass(frozen=True, slots=True)
39
+ class RunDirectories:
40
+ """One private, immutable set of runner-owned directories."""
41
+
42
+ root: Path
43
+ input_dir: Path
44
+ checkpoint_dir: Path
45
+ temporary_dir: Path
46
+ retain_artifacts: bool
47
+ _base_dir: Path = field(repr=False, compare=False)
48
+
49
+ def cleanup(self) -> None:
50
+ """Remove the generated run root unless artifact retention is enabled."""
51
+ if self.retain_artifacts:
52
+ return
53
+ if not self.root.exists():
54
+ return
55
+
56
+ resolved_root = self.root.resolve(strict=True)
57
+ _require_direct_child(resolved_root, self._base_dir, "run root")
58
+ shutil.rmtree(resolved_root)
59
+
60
+
61
+ def create_run_directories(
62
+ *,
63
+ base_dir: str | os.PathLike[str] | None = None,
64
+ retain_artifacts: bool = False,
65
+ ) -> RunDirectories:
66
+ """Create a unique, isolated directory layout for one scenario run."""
67
+ if not isinstance(retain_artifacts, bool):
68
+ raise TypeError("retain_artifacts must be a bool.")
69
+ if retain_artifacts and base_dir is None:
70
+ raise ValueError("retain_artifacts=True requires an explicit base directory.")
71
+
72
+ if base_dir is None:
73
+ root = Path(tempfile.mkdtemp(prefix=_RUN_DIRECTORY_PREFIX)).resolve(strict=True)
74
+ resolved_base = root.parent.resolve(strict=True)
75
+ else:
76
+ resolved_base = _resolve_base_directory(base_dir)
77
+ root = Path(tempfile.mkdtemp(prefix=_RUN_DIRECTORY_PREFIX, dir=resolved_base)).resolve(
78
+ strict=True
79
+ )
80
+
81
+ _require_direct_child(root, resolved_base, "run root")
82
+
83
+ input_dir = root / "input"
84
+ checkpoint_dir = root / "checkpoint"
85
+ temporary_dir = root / "temporary"
86
+ try:
87
+ input_dir.mkdir()
88
+ input_dir = input_dir.resolve(strict=True)
89
+ _require_direct_child(input_dir, root, "input directory")
90
+ checkpoint_dir.mkdir()
91
+ checkpoint_dir = checkpoint_dir.resolve(strict=True)
92
+ _require_direct_child(checkpoint_dir, root, "checkpoint directory")
93
+ temporary_dir.mkdir()
94
+ temporary_dir = temporary_dir.resolve(strict=True)
95
+ _require_direct_child(temporary_dir, root, "temporary directory")
96
+ except BaseException:
97
+ shutil.rmtree(root)
98
+ raise
99
+
100
+ return RunDirectories(
101
+ root=root,
102
+ input_dir=input_dir,
103
+ checkpoint_dir=checkpoint_dir,
104
+ temporary_dir=temporary_dir,
105
+ retain_artifacts=retain_artifacts,
106
+ _base_dir=resolved_base,
107
+ )
@@ -0,0 +1,59 @@
1
+ """Private atomic input-file publication for scenario batches."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ from dataclasses import dataclass, field
7
+ from pathlib import Path
8
+
9
+ from streamcase._directories import RunDirectories, _require_direct_child
10
+ from streamcase._serialization import _encode_batch_json_lines
11
+ from streamcase.actions import Batch
12
+
13
+ _BATCH_INDEX_WIDTH = 20
14
+
15
+
16
+ def _batch_file_name(index: int) -> str:
17
+ return f"batch-{index:0{_BATCH_INDEX_WIDTH}d}.json"
18
+
19
+
20
+ @dataclass(slots=True)
21
+ class AtomicBatchWriter:
22
+ """Publish encoded batches sequentially into one isolated input directory."""
23
+
24
+ directories: RunDirectories
25
+ _next_index: int = field(default=0, init=False, repr=False)
26
+
27
+ def publish(self, action: Batch) -> Path:
28
+ """Atomically publish one batch and return its final input path."""
29
+ encoded = _encode_batch_json_lines(action)
30
+ file_name = _batch_file_name(self._next_index)
31
+ destination = self.directories.input_dir / file_name
32
+ temporary = self.directories.input_dir / f".{file_name}.tmp"
33
+
34
+ _require_direct_child(destination, self.directories.input_dir, "batch input file")
35
+ _require_direct_child(temporary, self.directories.input_dir, "temporary batch file")
36
+
37
+ if destination.exists():
38
+ raise FileExistsError(f"Batch input destination already exists: {destination}")
39
+ if temporary.exists():
40
+ raise FileExistsError(f"Temporary batch input file already exists: {temporary}")
41
+
42
+ temporary_created = False
43
+ try:
44
+ with temporary.open("x", encoding="utf-8", newline="\n") as handle:
45
+ temporary_created = True
46
+ handle.write(encoded)
47
+ handle.flush()
48
+ os.fsync(handle.fileno())
49
+
50
+ if destination.exists():
51
+ raise FileExistsError(f"Batch input destination already exists: {destination}")
52
+ temporary.rename(destination)
53
+ except BaseException:
54
+ if temporary_created:
55
+ temporary.unlink(missing_ok=True)
56
+ raise
57
+
58
+ self._next_index += 1
59
+ return destination
@@ -0,0 +1,31 @@
1
+ """Internal serialization helpers for scenario actions."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+
7
+ from streamcase.actions import Batch
8
+
9
+
10
+ def _encode_batch_json_lines(action: Batch) -> str:
11
+ """Encode a batch as deterministic, newline-terminated JSON Lines text."""
12
+ lines: list[str] = []
13
+ for row_index, row in enumerate(action.rows):
14
+ try:
15
+ encoded = json.dumps(
16
+ dict(row),
17
+ allow_nan=False,
18
+ ensure_ascii=False,
19
+ separators=(",", ":"),
20
+ sort_keys=True,
21
+ )
22
+ except TypeError as error:
23
+ message = f"Batch row at index {row_index} could not be encoded as JSON: {error}"
24
+ raise TypeError(message) from error
25
+ except ValueError as error:
26
+ message = f"Batch row at index {row_index} could not be encoded as JSON: {error}"
27
+ raise ValueError(message) from error
28
+
29
+ lines.append(encoded)
30
+
31
+ return "\n".join(lines) + "\n"
@@ -0,0 +1,33 @@
1
+ """Private driver-side capture for Spark ``foreachBatch`` callbacks."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from _thread import LockType
6
+ from dataclasses import dataclass, field
7
+ from threading import Lock
8
+ from typing import TYPE_CHECKING
9
+
10
+ from streamcase.results import CapturedBatch
11
+
12
+ if TYPE_CHECKING:
13
+ from pyspark.sql import DataFrame
14
+
15
+
16
+ @dataclass(slots=True)
17
+ class _BatchCapture:
18
+ """Collect complete immutable output batches in callback order."""
19
+
20
+ _batches: list[CapturedBatch] = field(default_factory=list, init=False, repr=False)
21
+ _lock: LockType = field(default_factory=Lock, init=False, repr=False)
22
+
23
+ def callback(self, dataframe: DataFrame, batch_id: int) -> None:
24
+ """Collect and freeze one Spark output micro-batch on the driver."""
25
+ with self._lock:
26
+ rows = (row.asDict(recursive=True) for row in dataframe.collect())
27
+ captured = CapturedBatch(batch_id, rows)
28
+ self._batches.append(captured)
29
+
30
+ def snapshot(self) -> tuple[CapturedBatch, ...]:
31
+ """Return an immutable snapshot after any active callback completes."""
32
+ with self._lock:
33
+ return tuple(self._batches)
@@ -0,0 +1,95 @@
1
+ """Private execution of Spark streaming scenario actions."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Callable, Mapping
6
+ from typing import TYPE_CHECKING
7
+
8
+ from streamcase._cleanup import _cleanup_on_exit, _record_cleanup_failure
9
+ from streamcase._directories import RunDirectories
10
+ from streamcase._input_files import AtomicBatchWriter
11
+ from streamcase._spark_capture import _BatchCapture
12
+ from streamcase._spark_lifecycle import _QueryLifecycle
13
+ from streamcase.actions import Batch, Restart
14
+ from streamcase.results import ScenarioResult
15
+ from streamcase.scenario import Scenario
16
+
17
+ if TYPE_CHECKING:
18
+ from pyspark.sql import DataFrame
19
+
20
+
21
+ def _process_batch_action(
22
+ index: int,
23
+ action: Batch,
24
+ input_writer: AtomicBatchWriter,
25
+ lifecycle: _QueryLifecycle,
26
+ ) -> None:
27
+ try:
28
+ query = lifecycle.require_active()
29
+ input_writer.publish(action)
30
+ query.processAllAvailable()
31
+ lifecycle.require_active()
32
+ except Exception as error:
33
+ raise RuntimeError(f"Batch action at index {index} failed: {error}") from error
34
+
35
+
36
+ def _process_restart_action(
37
+ index: int,
38
+ rebuild_stream: Callable[[], DataFrame],
39
+ lifecycle: _QueryLifecycle,
40
+ ) -> None:
41
+ try:
42
+ lifecycle.require_active()
43
+ lifecycle.stop()
44
+ except Exception as error:
45
+ raise _restart_failure(index, "stopping the active query", error) from error
46
+
47
+ try:
48
+ replacement_stream = rebuild_stream()
49
+ except Exception as error:
50
+ raise _restart_failure(index, "rebuilding the stream", error) from error
51
+
52
+ try:
53
+ lifecycle.start(replacement_stream, action_index=index + 1)
54
+ lifecycle.require_active()
55
+ except Exception as error:
56
+ raise _restart_failure(index, "starting the replacement query", error) from error
57
+
58
+
59
+ def _restart_failure(index: int, transition: str, error: Exception) -> RuntimeError:
60
+ failure = RuntimeError(f"Restart action at index {index} failed while {transition}: {error}")
61
+ for resource, cleanup_error in getattr(error, "_streamcase_cleanup_failures", ()):
62
+ _record_cleanup_failure(failure, resource, cleanup_error)
63
+ return failure
64
+
65
+
66
+ def _execute_batches(
67
+ stream: DataFrame,
68
+ scenario: Scenario,
69
+ directories: RunDirectories,
70
+ capture: _BatchCapture,
71
+ *,
72
+ output_mode: str = "append",
73
+ query_options: Mapping[str, str] | None = None,
74
+ rebuild_stream: Callable[[], DataFrame] | None = None,
75
+ ) -> ScenarioResult:
76
+ """Execute each action and return an immutable output snapshot."""
77
+ if rebuild_stream is None and any(isinstance(action, Restart) for action in scenario.actions):
78
+ raise ValueError("Restart actions require a stream rebuild callback.")
79
+
80
+ lifecycle = _QueryLifecycle(
81
+ directories,
82
+ capture,
83
+ output_mode=output_mode,
84
+ query_options=query_options,
85
+ )
86
+ with _cleanup_on_exit("streaming query", lifecycle.stop):
87
+ lifecycle.start(stream, action_index=0)
88
+ input_writer = AtomicBatchWriter(directories)
89
+ for index, action in enumerate(scenario.actions):
90
+ if isinstance(action, Batch):
91
+ _process_batch_action(index, action, input_writer, lifecycle)
92
+ else:
93
+ assert rebuild_stream is not None
94
+ _process_restart_action(index, rebuild_stream, lifecycle)
95
+ return ScenarioResult(capture.snapshot())
@@ -0,0 +1,81 @@
1
+ """Private ownership and transitions for one Spark streaming query at a time."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Mapping
6
+ from typing import TYPE_CHECKING
7
+
8
+ from streamcase._cleanup import _cleanup_on_exit
9
+ from streamcase._directories import RunDirectories
10
+ from streamcase._spark_capture import _BatchCapture
11
+ from streamcase._spark_query_options import _prepare_query_configuration
12
+
13
+ if TYPE_CHECKING:
14
+ from pyspark.sql import DataFrame
15
+ from pyspark.sql.streaming import StreamingQuery
16
+
17
+
18
+ def _stop_named_query(stream: DataFrame, name: str) -> None:
19
+ for active_query in stream.sparkSession.streams.active:
20
+ if active_query.name == name:
21
+ active_query.stop()
22
+
23
+
24
+ class _QueryLifecycle:
25
+ """Keep runner-owned query transitions and writer settings in one place."""
26
+
27
+ def __init__(
28
+ self,
29
+ directories: RunDirectories,
30
+ capture: _BatchCapture,
31
+ *,
32
+ output_mode: str = "append",
33
+ query_options: Mapping[str, str] | None = None,
34
+ ) -> None:
35
+ self._directories = directories
36
+ self._capture = capture
37
+ self._output_mode, self._query_options = _prepare_query_configuration(
38
+ output_mode, query_options
39
+ )
40
+ self._query: StreamingQuery | None = None
41
+ self._pending_start_stream: DataFrame | None = None
42
+
43
+ def require_active(self) -> StreamingQuery:
44
+ """Return the owned active query, rejecting missing or terminated state."""
45
+ if self._query is None:
46
+ raise RuntimeError("No streaming query is owned by this run.")
47
+ if not self._query.isActive:
48
+ raise RuntimeError("The streaming query terminated.")
49
+ return self._query
50
+
51
+ def start(self, stream: DataFrame, *, action_index: int) -> None:
52
+ """Start a query only when the previous one was successfully stopped."""
53
+ if self._query is not None or self._pending_start_stream is not None:
54
+ raise RuntimeError("Stop the owned streaming query before starting another.")
55
+
56
+ writer = stream.writeStream
57
+ for name, value in self._query_options.items():
58
+ writer = writer.option(name, value)
59
+ writer = (
60
+ writer.foreachBatch(self._capture.callback)
61
+ .outputMode(self._output_mode)
62
+ .option("checkpointLocation", str(self._directories.checkpoint_dir))
63
+ .queryName(self._directories.root.name)
64
+ )
65
+ try:
66
+ self._query = writer.start()
67
+ except Exception as error:
68
+ self._pending_start_stream = stream
69
+ with _cleanup_on_exit("streaming query", self.stop):
70
+ raise RuntimeError(
71
+ f"Could not start the query before Batch action at index {action_index}."
72
+ ) from error
73
+
74
+ def stop(self) -> None:
75
+ """Stop once; repeat calls are safe after success or before first start."""
76
+ if self._query is not None:
77
+ self._query.stop()
78
+ self._query = None
79
+ if self._pending_start_stream is not None:
80
+ _stop_named_query(self._pending_start_stream, self._directories.root.name)
81
+ self._pending_start_stream = None
@@ -0,0 +1,43 @@
1
+ """Private ownership boundary for one Spark scenario run."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ from collections.abc import Callable, Mapping
7
+ from typing import TYPE_CHECKING
8
+
9
+ from streamcase._cleanup import _cleanup_on_exit
10
+ from streamcase._directories import RunDirectories, create_run_directories
11
+ from streamcase._spark_capture import _BatchCapture
12
+ from streamcase._spark_execution import _execute_batches
13
+ from streamcase._spark_query_options import _prepare_query_configuration
14
+ from streamcase.results import ScenarioResult
15
+ from streamcase.scenario import Scenario
16
+
17
+ if TYPE_CHECKING:
18
+ from pyspark.sql import DataFrame
19
+
20
+
21
+ def _run_managed_batches(
22
+ scenario: Scenario,
23
+ build_stream: Callable[[RunDirectories], DataFrame],
24
+ *,
25
+ base_dir: str | os.PathLike[str] | None = None,
26
+ retain_artifacts: bool = False,
27
+ output_mode: str = "append",
28
+ query_options: Mapping[str, str] | None = None,
29
+ ) -> ScenarioResult:
30
+ """Own run directories while a caller-owned session builds the stream."""
31
+ approved_mode, approved_options = _prepare_query_configuration(output_mode, query_options)
32
+ directories = create_run_directories(base_dir=base_dir, retain_artifacts=retain_artifacts)
33
+ with _cleanup_on_exit("run directory", directories.cleanup):
34
+ stream = build_stream(directories)
35
+ return _execute_batches(
36
+ stream,
37
+ scenario,
38
+ directories,
39
+ _BatchCapture(),
40
+ output_mode=approved_mode,
41
+ query_options=approved_options,
42
+ rebuild_stream=lambda: build_stream(directories),
43
+ )
@@ -0,0 +1,67 @@
1
+ """Validate the narrow, runner-safe Spark streaming writer configuration."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Mapping
6
+
7
+ _OUTPUT_MODES = frozenset({"append", "complete", "update"})
8
+ _RESERVED_QUERY_OPTIONS = frozenset(
9
+ {
10
+ "availablenow",
11
+ "checkpointlocation",
12
+ "clusterby",
13
+ "continuous",
14
+ "foreach",
15
+ "foreachbatch",
16
+ "format",
17
+ "once",
18
+ "outputmode",
19
+ "partitionby",
20
+ "path",
21
+ "paths",
22
+ "processingtime",
23
+ "queryname",
24
+ "sink",
25
+ "table",
26
+ "totable",
27
+ "trigger",
28
+ }
29
+ )
30
+
31
+
32
+ def _prepare_query_configuration(
33
+ output_mode: str,
34
+ query_options: Mapping[str, str] | None,
35
+ ) -> tuple[str, dict[str, str]]:
36
+ """Return a validated mode and detached copy of caller writer options."""
37
+ if not isinstance(output_mode, str):
38
+ raise TypeError("Spark output mode must be a string.")
39
+ if output_mode not in _OUTPUT_MODES:
40
+ raise ValueError("Spark output mode must be 'append', 'complete', or 'update'.")
41
+
42
+ if query_options is None:
43
+ return output_mode, {}
44
+ if not isinstance(query_options, Mapping):
45
+ raise TypeError("Spark query options must be a mapping of strings to strings.")
46
+
47
+ copied = dict(query_options)
48
+ seen: set[str] = set()
49
+ conflicts: list[str] = []
50
+ for name, value in copied.items():
51
+ if not isinstance(name, str):
52
+ raise TypeError("Spark query option names must be strings.")
53
+ if not isinstance(value, str):
54
+ raise TypeError(f"Spark query option {name!r} must have a string value.")
55
+
56
+ normalized = name.casefold()
57
+ if normalized in seen:
58
+ raise ValueError(f"Spark query option {name!r} duplicates a case-insensitive name.")
59
+ seen.add(normalized)
60
+ if normalized in _RESERVED_QUERY_OPTIONS:
61
+ conflicts.append(name)
62
+
63
+ if conflicts:
64
+ rendered = ", ".join(repr(name) for name in sorted(conflicts, key=str.casefold))
65
+ raise ValueError(f"Runner-owned query options cannot be overridden: {rendered}.")
66
+
67
+ return output_mode, copied