onedata-lambda-sdk 1.0.0__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.
@@ -0,0 +1,388 @@
1
+ """
2
+ The lambda runtime contract: `run()` -- the single entry point the base image's
3
+ bootstrap calls (`from onedata_lambda_sdk import run; run(handler, raw)`).
4
+
5
+ The logic is migrated from the v2 base `index.py` (which ran as a fresh process per
6
+ request); here it lives in the SDK and is driven in-process. The `@per_job` adapter that
7
+ produces a batch handler for `run()` lives in `perjob`.
8
+ """
9
+
10
+ __author__ = "Bartosz Walkowicz"
11
+ __copyright__ = "Copyright (C) 2026 Onedata (onedata.org)"
12
+ __license__ = "This software is released under the MIT license cited in LICENSE.txt"
13
+
14
+ import json
15
+ import os
16
+ import sys
17
+ import threading
18
+ import time
19
+ import traceback
20
+ import uuid
21
+ from collections.abc import Callable, Iterator
22
+ from contextlib import contextmanager
23
+ from datetime import datetime
24
+ from typing import Any, Final
25
+
26
+ import requests
27
+
28
+ from ._wire import AtmJobArgs, AtmJobBatchRequest, AtmJobBatchRequestCtx, AtmJobBatchResponse
29
+ from .job import Job, JobContext, JobException
30
+ from .logging import DEFAULT_LOG_LEVEL, Logger
31
+ from .oneclient import ENV_MOUNT_POINT
32
+ from .streaming import ResultStreamer, StreamFlusher
33
+ from .types import AtmException
34
+
35
+
36
+ # --- Environment contract --------------------------------------------------------------
37
+ # Every environment variable this module reads, named in one place so the runtime's
38
+ # dependency on the environment is visible at a glance. Read via `_get_env` at call time
39
+ # (not import time) so tests and dev-mode can set them per run.
40
+ ENV_VERIFY_SSL: Final[str] = "VERIFY_SSL_CERTIFICATES"
41
+ ENV_DEBUG_MODE: Final[str] = "OPENFAAS_FUNCTION_DEBUG_MODE"
42
+ ENV_ONECLIENT_MOUNTED: Final[str] = "ONECLIENT_MOUNTED"
43
+ ENV_DEV_MODE: Final[str] = "ONEDATA_LAMBDA_DEV"
44
+ # The mount-point env name is owned by `oneclient` (one source of truth shared with the
45
+ # `mounted_file_path` helper handlers use), re-exported here for the mount-wait below.
46
+
47
+ # Local file the runtime's diagnostics (and any stray handler prints) are redirected to,
48
+ # so they never reach stdout/stderr (which OpenFaaS reads as the response).
49
+ LAMBDA_OUTPUT_LOG: Final[str] = "/tmp/lambda.out"
50
+
51
+ HEARTBEAT_TIMEOUT_SEC: Final[int] = 10
52
+ FIRST_HEARTBEAT_RETRIES: Final[int] = 10
53
+ FIRST_HEARTBEAT_RETRY_DELAY_SEC: Final[int] = 5
54
+
55
+ MOUNT_POINT_TIMEOUT_SEC: Final[int] = 120
56
+ MOUNT_POINT_POLL_INTERVAL_SEC: Final[float] = 0.5
57
+
58
+ #: A runtime diagnostic logger: `log(severity, message)` -> /tmp/lambda.out.
59
+ DiagLogger = Callable[[str, str], None]
60
+
61
+
62
+ class _FirstHeartbeatError(Exception):
63
+ pass
64
+
65
+
66
+ class _OneclientMountError(Exception):
67
+ pass
68
+
69
+
70
+ def run(
71
+ handler: Callable[..., Any], raw: str, *, out_dir: str | None = None
72
+ ) -> AtmJobBatchResponse[Any] | AtmException:
73
+ """
74
+ Run a batch-level `handle(jobs, ctx)` handler against a raw OpenFaaS request and
75
+ return the response envelope as a dict.
76
+
77
+ Owns everything between the raw request and the handler:
78
+
79
+ * parse the request into the per-batch `JobContext` and the `Job` list;
80
+ * redirect stdout/stderr to LAMBDA_OUTPUT_LOG so stray handler prints don't corrupt the
81
+ response, and write the runtime's own lifecycle diagnostics there;
82
+ * deliver the mandatory first heartbeat, then vend a throttled `heartbeat()` service;
83
+ * wait for the Oneclient mount when the lambda uses one;
84
+ * dispatch the handler and assemble the envelope.
85
+
86
+ Never raises for an expected failure: a whole-batch handler error (or a malformed
87
+ request) becomes a top-level `{"exception": ...}`; per-job failures are the
88
+ `AtmException` entries the handler placed in its `{"resultsBatch": ...}`.
89
+
90
+ `out_dir` overrides where result/log streams are written (default `/out`); production
91
+ leaves it unset, tests inject a temp dir (see `testing.run_local`).
92
+ """
93
+ batch_id = uuid.uuid4()
94
+ debug_mode = _get_env(ENV_DEBUG_MODE) == "true"
95
+
96
+ with _redirected_streams():
97
+
98
+ def log(severity: str, message: str) -> None:
99
+ _diag_log(batch_id, severity, message)
100
+
101
+ log("info", "Started request processing")
102
+ if debug_mode:
103
+ log("debug", "Lambda is running in debug mode")
104
+
105
+ try:
106
+ _await_oneclient_mount(log)
107
+
108
+ # Shape assertion, not validation: json.loads returns Any, and malformed input
109
+ # surfaces as the top-level exception envelope below rather than a type error.
110
+ request: AtmJobBatchRequest[Any] = json.loads(raw)
111
+ if debug_mode:
112
+ log("debug", f"Input: {request}")
113
+
114
+ ctx_wire = request["ctx"]
115
+
116
+ send_heartbeat = _make_heartbeat_sender(ctx_wire, log)
117
+ # The first heartbeat is mandatory and unconditional: it tells Oneprovider
118
+ # the job batch actually started executing.
119
+ log("info", "Sending first heartbeat")
120
+ _deliver_first_heartbeat(send_heartbeat)
121
+
122
+ # Narrow the internal sender to the public heartbeat service: the inward API
123
+ # (JobContext.heartbeat, the streamers, the flusher) is typed `Callable[[], None]`,
124
+ # while `send_heartbeat` is `(*, force: bool) -> bool`. This adapter hides the
125
+ # `force=` throttle-bypass (runtime-only, for the first heartbeat) and the success
126
+ # bool, so neither leaks into the inward-facing contract.
127
+ def heartbeat() -> None:
128
+ send_heartbeat()
129
+
130
+ flusher = StreamFlusher(heartbeat=heartbeat, out_dir=out_dir)
131
+ try:
132
+ jobs, job_ctx = _build_jobs_and_context(request, heartbeat, flusher, out_dir)
133
+
134
+ log("info", "Calling handler callback")
135
+ results = handler(jobs, job_ctx)
136
+ response: AtmJobBatchResponse[Any] = {"resultsBatch": results}
137
+ if debug_mode:
138
+ log("debug", f"Output: {response}")
139
+ log("info", "Ended request processing")
140
+
141
+ return response
142
+ finally:
143
+ # Drain buffered streams/logs before responding, even on error.
144
+ flusher.stop()
145
+ except _OneclientMountError:
146
+ log("error", "Request processing failed to mount Oneclient")
147
+ return AtmException(exception="Failed to mount Oneclient")
148
+ except _FirstHeartbeatError:
149
+ log("error", "Request processing failed to deliver first heartbeat")
150
+ return AtmException(exception="Failed to deliver first heartbeat")
151
+ except JobException as ex:
152
+ # A known, whole-batch failure raised by a (batch-style) handler -- report
153
+ # just the message, no traceback noise. Per-job failures never reach here:
154
+ # `@per_job` turns them into AtmException entries inside the results batch.
155
+ log("error", f"Request processing failed: {ex}")
156
+ return AtmException(exception=str(ex))
157
+ except Exception as ex:
158
+ log("error", f"Request processing failed with unexpected exception: {ex}")
159
+ return AtmException(exception=traceback.format_exc())
160
+
161
+
162
+ @contextmanager
163
+ def _redirected_streams() -> Iterator[None]:
164
+ """
165
+ Redirect fd 1/2 to LAMBDA_OUTPUT_LOG for the duration of the call.
166
+
167
+ OpenFaaS treats anything written to the function's stdout as (part of) the response,
168
+ so accidental `print`s from a handler -- or a library it uses -- would corrupt it.
169
+ We redirect at the file-descriptor level (not just `sys.stdout`) so writes from C
170
+ extensions and subprocesses are captured too, and restore on the way out so the
171
+ bootstrap can write the real response. The runtime's own diagnostics go to stderr,
172
+ which is redirected here too, so they land in the same file.
173
+ """
174
+ if _is_dev_mode():
175
+ # Keep stdout/stderr so handler output and diagnostics are visible locally.
176
+ yield
177
+ return
178
+
179
+ sink = open(LAMBDA_OUTPUT_LOG, "a") # noqa: SIM115 -- manual fd lifecycle (restored in finally)
180
+ saved_stdout_fd = os.dup(sys.stdout.fileno())
181
+ saved_stderr_fd = os.dup(sys.stderr.fileno())
182
+ try:
183
+ os.dup2(sink.fileno(), sys.stdout.fileno()) # $ exec >&tmp_out_fd
184
+ os.dup2(sink.fileno(), sys.stderr.fileno()) # $ exec 2>&tmp_out_fd
185
+ yield
186
+ finally:
187
+ sys.stdout.flush()
188
+ sys.stderr.flush()
189
+ os.dup2(saved_stdout_fd, sys.stdout.fileno())
190
+ os.dup2(saved_stderr_fd, sys.stderr.fileno())
191
+ # `os.dup` above allocated two fresh descriptors; once dup2 has restored fd 1/2 to
192
+ # point at them, the saved copies are redundant. Close them (and the sink) or we
193
+ # leak two fds per request -- and run() is called once per request in a long-lived
194
+ # process, so the leak would accumulate until the fd limit is hit. Closing the
195
+ # copies does not disturb the just-restored fd 1/2 (independent descriptors).
196
+ os.close(saved_stdout_fd)
197
+ os.close(saved_stderr_fd)
198
+ sink.close()
199
+
200
+
201
+ def _diag_log(job_id: uuid.UUID, severity: str, message: str) -> None:
202
+ # The runtime's own lifecycle trace. In production stdout/stderr are redirected to
203
+ # LAMBDA_OUTPUT_LOG (see _redirected_streams), so this lands there; in dev mode it
204
+ # stays on the console. Distinct from the audit-log stream a handler writes via
205
+ # ctx.logger(name).
206
+ timestamp = datetime.now().strftime("%m/%d/%Y, %H:%M:%S:%f")
207
+ print(f"[{severity.upper()} {timestamp} {job_id}] {message}", file=sys.stderr)
208
+
209
+
210
+ def _await_oneclient_mount(log: DiagLogger) -> None:
211
+ if _is_dev_mode():
212
+ log("debug", "Skipping Oneclient mount wait (dev mode)")
213
+ return
214
+
215
+ if _get_env(ENV_ONECLIENT_MOUNTED) != "true":
216
+ log("info", "This lambda is not using Oneclient")
217
+ return
218
+
219
+ mount_point = _get_env(ENV_MOUNT_POINT)
220
+ if not mount_point:
221
+ log("error", "Oneclient declared mounted but no mount point given")
222
+ raise _OneclientMountError
223
+
224
+ log("info", f"This lambda is using Oneclient; mount point: {mount_point}")
225
+ deadline = time.monotonic() + MOUNT_POINT_TIMEOUT_SEC
226
+ while time.monotonic() < deadline:
227
+ if any(os.listdir(mount_point)):
228
+ log("info", f"Oneclient mounted: {os.listdir(mount_point)}")
229
+ return
230
+ time.sleep(MOUNT_POINT_POLL_INTERVAL_SEC)
231
+
232
+ log("error", "Did not manage to mount Oneclient")
233
+ raise _OneclientMountError
234
+
235
+
236
+ def _make_heartbeat_sender(
237
+ ctx_wire: AtmJobBatchRequestCtx[Any], log: DiagLogger
238
+ ) -> Callable[..., bool]:
239
+ """
240
+ Build a throttled heartbeat sender returning whether a recent heartbeat is in place.
241
+
242
+ Throttled to one POST per `0.1 * timeoutSeconds` (a call within that window is a
243
+ no-op reported as success, since a fresh heartbeat already exists). `force=True`
244
+ bypasses the throttle (used for the mandatory first heartbeat).
245
+ """
246
+ if _is_dev_mode():
247
+
248
+ def log_heartbeat(*, force: bool = False) -> bool:
249
+ log("debug", f"Heartbeat (force={force}) [dev mode, not sent]")
250
+ return True
251
+
252
+ return log_heartbeat
253
+
254
+ url = ctx_wire["heartbeatUrl"]
255
+ throttle_period = 0.1 * ctx_wire["timeoutSeconds"]
256
+ verify = _get_env(ENV_VERIFY_SSL) == "true"
257
+ last_sent_at: float | None = None
258
+ # Guards last_sent_at so parallel `@per_job` workers don't tear it. The POST runs
259
+ # outside the lock (no serializing the network); the worst a race can do is send one
260
+ # extra heartbeat, which is harmless.
261
+ lock = threading.Lock()
262
+
263
+ def send(*, force: bool = False) -> bool:
264
+ nonlocal last_sent_at
265
+ now = time.monotonic()
266
+ with lock:
267
+ recent = last_sent_at is not None and now - last_sent_at < throttle_period
268
+
269
+ if not force and recent:
270
+ return True
271
+
272
+ try:
273
+ response = requests.post(url, data={}, timeout=HEARTBEAT_TIMEOUT_SEC, verify=verify)
274
+ except requests.RequestException as ex:
275
+ log("warning", f"Failed to send heartbeat: {ex}")
276
+ return False
277
+
278
+ if response.ok:
279
+ with lock:
280
+ last_sent_at = now
281
+ log("info", "Heartbeat sent")
282
+ return True
283
+
284
+ log(
285
+ "warning",
286
+ f"Failed to send heartbeat due to:\n"
287
+ f"> response code: {response.status_code}\n"
288
+ f"> response content: {response.content!r}",
289
+ )
290
+ return False
291
+
292
+ return send
293
+
294
+
295
+ def _deliver_first_heartbeat(send_heartbeat: Callable[..., bool]) -> None:
296
+ for _ in range(FIRST_HEARTBEAT_RETRIES):
297
+ if send_heartbeat(force=True):
298
+ return
299
+ time.sleep(FIRST_HEARTBEAT_RETRY_DELAY_SEC)
300
+ raise _FirstHeartbeatError
301
+
302
+
303
+ def _build_jobs_and_context(
304
+ request: AtmJobBatchRequest[Any],
305
+ heartbeat: Callable[[], None],
306
+ flusher: StreamFlusher,
307
+ out_dir: str | None,
308
+ ) -> tuple[list[Job[Any]], JobContext[Any]]:
309
+ ctx_wire = request["ctx"]
310
+ log_level = ctx_wire.get("logLevel", DEFAULT_LOG_LEVEL)
311
+ jobs = [_build_job(item) for item in request["argsBatch"]]
312
+
313
+ # One streamer/logger instance per stream name, shared across the batch's jobs (they
314
+ # append to the same /out/<name> through the shared flusher).
315
+ streamers: dict[str, ResultStreamer[Any]] = {}
316
+
317
+ def result_streamer_factory(name: str, *, buffered: bool = True) -> ResultStreamer[Any]:
318
+ streamer = streamers.get(name)
319
+ if streamer is None:
320
+ streamer = ResultStreamer(
321
+ result_name=name,
322
+ buffered=buffered,
323
+ flusher=flusher,
324
+ heartbeat=heartbeat,
325
+ out_dir=out_dir,
326
+ )
327
+ streamers[name] = streamer
328
+ return streamer
329
+
330
+ loggers: dict[str, Logger] = {}
331
+
332
+ def logger_factory(name: str) -> Logger:
333
+ # Logs are low-frequency and worth keeping even on a crash -> direct (unbuffered).
334
+ logger = loggers.get(name)
335
+ if logger is None:
336
+ logger = Logger(
337
+ result_name=name,
338
+ buffered=False,
339
+ flusher=flusher,
340
+ heartbeat=heartbeat,
341
+ log_level=log_level,
342
+ out_dir=out_dir,
343
+ )
344
+ loggers[name] = logger
345
+ return logger
346
+
347
+ job_ctx: JobContext[Any] = JobContext(
348
+ config=ctx_wire.get("config"),
349
+ user_id=ctx_wire["userId"],
350
+ space_id=ctx_wire["spaceId"],
351
+ workflow_execution_id=ctx_wire["atmWorkflowExecutionId"],
352
+ oneprovider_id=ctx_wire["oneproviderId"],
353
+ oneprovider_domain=ctx_wire["oneproviderDomain"],
354
+ onezone_domain=ctx_wire["onezoneDomain"],
355
+ access_token=ctx_wire["accessToken"],
356
+ timeout_seconds=ctx_wire["timeoutSeconds"],
357
+ log_level=log_level,
358
+ heartbeat=heartbeat,
359
+ result_streamer_factory=result_streamer_factory,
360
+ logger_factory=logger_factory,
361
+ raw=ctx_wire,
362
+ )
363
+ return jobs, job_ctx
364
+
365
+
366
+ def _build_job(item: AtmJobArgs) -> Job[Any]:
367
+ # trace_id comes from the wire: the framework assigns it per job in a `__meta`
368
+ # envelope. We never fabricate it -- a missing __meta.traceId is a malformed request
369
+ # (top-level exception), so contract drift surfaces instead of being masked. The
370
+ # envelope is stripped so the handler sees only its declared arguments.
371
+ try:
372
+ trace_id = item["__meta"]["traceId"]
373
+ except (KeyError, TypeError) as ex:
374
+ raise RuntimeError(f"job args missing __meta.traceId: {ex}") from ex
375
+
376
+ args = {key: value for key, value in item.items() if key != "__meta"}
377
+ return Job(args=args, trace_id=trace_id)
378
+
379
+
380
+ def _get_env(name: str, default: str | None = None) -> str | None:
381
+ return os.environ.get(name, default)
382
+
383
+
384
+ def _is_dev_mode() -> bool:
385
+ # Local dev / runtime-test mode (ONEDATA_LAMBDA_DEV=1): log heartbeats instead of
386
+ # POSTing them, skip the stdout/stderr redirect, and skip the Oneclient mount-wait,
387
+ # so the real runtime can be driven locally on a fixture. See the `testing` module.
388
+ return _get_env(ENV_DEV_MODE) == "1"
@@ -0,0 +1,53 @@
1
+ """Time-series measurements for lambdas."""
2
+
3
+ __author__ = "Bartosz Walkowicz"
4
+ __copyright__ = "Copyright (C) 2022-2026 Onedata (onedata.org)"
5
+ __license__ = "This software is released under the MIT license cited in LICENSE.txt"
6
+
7
+
8
+ import time
9
+
10
+ from .types import AtmTimeSeriesMeasurement
11
+
12
+
13
+ class TimeSeriesMeasurementBuilder:
14
+ """Base class for time series measurement builders.
15
+
16
+ In order to create a builder for concrete time series measurement, create
17
+ a class deriving from TimeSeriesMeasurementBuilder with following
18
+ metadata specified:
19
+ - ts_name - used as exact name of the measurement
20
+ - unit - stored only as metadata and not processed in any way
21
+
22
+ Example usage:
23
+ >>>
24
+ >>> import pprint
25
+ >>>
26
+ >>> from onedata_lambda_sdk import TimeSeriesMeasurementBuilder
27
+ >>>
28
+ >>>
29
+ >>> class FilesProcessed(
30
+ >>> TimeSeriesMeasurementBuilder, ts_name="filesProcessed", unit=None
31
+ >>> ):
32
+ >>> pass
33
+ >>>
34
+ >>>
35
+ >>> pprint.pprint(FilesProcessed.build(value=34, timestamp=100))
36
+
37
+ {'timestamp': 100, 'tsName': 'filesProcessed', 'value': 34}
38
+ """
39
+
40
+ _ts_name: str
41
+ _unit: str | None
42
+
43
+ def __init_subclass__(cls, ts_name: str, unit: str | None) -> None:
44
+ cls._ts_name = ts_name
45
+ cls._unit = unit
46
+
47
+ @classmethod
48
+ def build(cls, value: float, *, timestamp: int | None = None) -> AtmTimeSeriesMeasurement:
49
+ return {
50
+ "tsName": cls._ts_name,
51
+ "timestamp": int(time.time()) if timestamp is None else timestamp,
52
+ "value": value,
53
+ }
@@ -0,0 +1,158 @@
1
+ """
2
+ Streaming lambda results and measurements to `/out/<name>`.
3
+
4
+ Beyond the per-job results a handler returns, a lambda can append items incrementally to
5
+ named streams (audit logs, time-series measurements, large result sets); this module
6
+ owns that path. The runtime owns the flusher's lifecycle -- lazily started on first
7
+ buffered use and drained-and-stopped when the handler returns (see `runtime.run`).
8
+ Everything here is thread-safe.
9
+ """
10
+
11
+ __author__ = "Bartosz Walkowicz"
12
+ __copyright__ = "Copyright (C) 2022-2026 Onedata (onedata.org)"
13
+ __license__ = "This software is released under the MIT license cited in LICENSE.txt"
14
+
15
+ import json
16
+ import os
17
+ import threading
18
+ from collections.abc import Callable, Iterable
19
+ from typing import Any, Final
20
+
21
+
22
+ OUT_DIR: Final[str] = "/out"
23
+
24
+ # How often the background flusher drains buffered items.
25
+ FLUSH_INTERVAL_SEC: Final[float] = 1.0
26
+
27
+
28
+ class StreamFlusher:
29
+ """
30
+ A single background thread batching writes for all buffered streams of one run.
31
+
32
+ Producers `enqueue` `(stream_name, item)` pairs (cheap, lock-guarded). The
33
+ thread wakes periodically, writes everything buffered grouped by stream, and -- only
34
+ if it actually wrote something -- emits a heartbeat (itself throttled by the
35
+ runtime). The thread starts lazily on first `enqueue` so idle lambdas pay nothing.
36
+ """
37
+
38
+ def __init__(
39
+ self,
40
+ *,
41
+ heartbeat: Callable[[], None],
42
+ out_dir: str | None = None,
43
+ interval: float = FLUSH_INTERVAL_SEC,
44
+ ) -> None:
45
+ self._heartbeat = heartbeat
46
+ self._out_dir = out_dir if out_dir is not None else OUT_DIR
47
+ self._interval = interval
48
+ self._buffer: list[tuple[str, Any]] = []
49
+ self._buffer_lock = threading.Lock()
50
+ self._start_lock = threading.Lock()
51
+ self._stop = threading.Event()
52
+ self._thread: threading.Thread | None = None
53
+
54
+ def enqueue(self, name: str, item: Any) -> None:
55
+ self.enqueue_many(name, [item])
56
+
57
+ def enqueue_many(self, name: str, items: list[Any]) -> None:
58
+ if not items:
59
+ return
60
+
61
+ self._ensure_started()
62
+ with self._buffer_lock:
63
+ self._buffer.extend((name, item) for item in items)
64
+
65
+ def _ensure_started(self) -> None:
66
+ if self._thread is not None:
67
+ return
68
+
69
+ with self._start_lock:
70
+ if self._thread is None:
71
+ thread = threading.Thread(target=self._loop, daemon=True)
72
+ self._thread = thread
73
+ thread.start()
74
+
75
+ def _loop(self) -> None:
76
+ while not self._stop.is_set():
77
+ self._stop.wait(self._interval)
78
+ self._flush_once()
79
+
80
+ def _flush_once(self) -> None:
81
+ with self._buffer_lock:
82
+ if not self._buffer:
83
+ return
84
+
85
+ pending = self._buffer
86
+ self._buffer = []
87
+
88
+ by_name: dict[str, list[Any]] = {}
89
+ for name, item in pending:
90
+ by_name.setdefault(name, []).append(item)
91
+ for name, items in by_name.items():
92
+ _append_to_stream(self._out_dir, name, items)
93
+
94
+ # Activity-gated: heartbeat only because we made observable progress.
95
+ self._heartbeat()
96
+
97
+ def stop(self) -> None:
98
+ """Stop the thread and drain remaining buffered items (no-op if unused)."""
99
+ self._stop.set()
100
+ if self._thread is not None:
101
+ self._thread.join()
102
+ self._flush_once()
103
+
104
+
105
+ class ResultStreamer[T]:
106
+ """
107
+ Streams items to `/out/<result_name>`. Two modes, neither needing a manual flush:
108
+
109
+ * **buffered (default)** -- enqueues into the shared background flusher, which batches
110
+ writes across all buffered streams and heartbeats after each non-empty flush. Best
111
+ for high-frequency / multi-threaded producers.
112
+ * **direct** (`buffered=False`) -- writes immediately and heartbeats (throttled) per
113
+ write. Best for low-frequency output or callers doing their own batching.
114
+
115
+ Vended by `JobContext.result_streamer(name, *, buffered=...)`; the runtime injects
116
+ the shared flusher and the heartbeat. Construct directly only in tests.
117
+ """
118
+
119
+ def __init__(
120
+ self,
121
+ *,
122
+ result_name: str,
123
+ buffered: bool,
124
+ flusher: StreamFlusher,
125
+ heartbeat: Callable[[], None],
126
+ out_dir: str | None = None,
127
+ ) -> None:
128
+ self.result_name = result_name
129
+ self._buffered = buffered
130
+ self._flusher = flusher
131
+ self._heartbeat = heartbeat
132
+ self._out_dir = out_dir if out_dir is not None else OUT_DIR
133
+ self._write_lock = threading.Lock()
134
+
135
+ def stream_item(self, item: T) -> None:
136
+ self.stream_items([item])
137
+
138
+ def stream_items(self, items: Iterable[T]) -> None:
139
+ materialized = list(items)
140
+ if not materialized:
141
+ return
142
+
143
+ if self._buffered:
144
+ # One lock acquisition for the whole batch, not one per item.
145
+ self._flusher.enqueue_many(self.result_name, materialized)
146
+ else:
147
+ with self._write_lock:
148
+ _append_to_stream(self._out_dir, self.result_name, materialized)
149
+
150
+ # Direct mode: heartbeat per write (throttled by the runtime).
151
+ self._heartbeat()
152
+
153
+
154
+ def _append_to_stream(out_dir: str, name: str, items: list[Any]) -> None:
155
+ with open(os.path.join(out_dir, name), "a") as stream_file:
156
+ for item in items:
157
+ json.dump(item, stream_file)
158
+ stream_file.write("\n")