lancedb-ray 0.1.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,47 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ """lancedb-ray: Ray Data integration for LanceDB and LanceDB Enterprise.
3
+
4
+ Read and write LanceDB tables as Ray Datasets, using the most parallel strategy
5
+ each backend supports:
6
+
7
+ - **Local / OSS** tables are backed by a Lance dataset, so reads run one task
8
+ per fragment group and appends write fragments in parallel that the driver
9
+ commits as a single atomic transaction.
10
+ - **Cloud / Enterprise** (``db://``) tables are a remote service with no
11
+ fragment access, so reads shard the row space across tasks against a pinned
12
+ table version and writes fan out batched requests.
13
+
14
+ Example:
15
+ >>> import lancedb_ray as ldbr
16
+ >>> ds = ldbr.read_lancedb("my_table", uri="/data/lancedb") # doctest: +SKIP
17
+ >>> ldbr.write_lancedb(ds, "copy", uri="/data/lancedb", mode="create") # doctest: +SKIP
18
+ """
19
+
20
+ from importlib.metadata import PackageNotFoundError
21
+ from importlib.metadata import version as _dist_version
22
+
23
+ from ._plan import OffsetRange
24
+ from ._retry import RetryPolicy
25
+ from .connection import LanceDBConnectionSpec
26
+ from .datasink import LanceDBDatasink, WriteStats
27
+ from .datasource import LanceDBDatasource
28
+ from .io import read_lancedb, write_lancedb
29
+
30
+ # Read from the installed distribution rather than hard-coded here, so the
31
+ # version release-please writes into pyproject.toml is the only one there is.
32
+ try:
33
+ __version__ = _dist_version("lancedb-ray")
34
+ except PackageNotFoundError: # pragma: no cover - only in an uninstalled tree
35
+ __version__ = "0.0.0.dev0"
36
+
37
+ __all__ = [
38
+ "LanceDBConnectionSpec",
39
+ "LanceDBDatasink",
40
+ "LanceDBDatasource",
41
+ "OffsetRange",
42
+ "RetryPolicy",
43
+ "WriteStats",
44
+ "__version__",
45
+ "read_lancedb",
46
+ "write_lancedb",
47
+ ]
lancedb_ray/_plan.py ADDED
@@ -0,0 +1,149 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ """Pure planning helpers for distributed LanceDB reads and writes.
3
+
4
+ Nothing in this module touches Ray, the network, or the filesystem. Keeping the
5
+ planning arithmetic free of I/O is what makes it cheap to test exhaustively --
6
+ the interesting edge cases (empty tables, parallelism exceeding row counts,
7
+ uneven remainders) are all covered by ``tests/test_plan.py``.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from collections.abc import Iterator
13
+ from typing import NamedTuple
14
+
15
+ __all__ = [
16
+ "OffsetRange",
17
+ "chunk_offsets",
18
+ "plan_offset_shards",
19
+ "rows_within_byte_budget",
20
+ "split_arrow_table",
21
+ ]
22
+
23
+
24
+ class OffsetRange(NamedTuple):
25
+ """A half-open ``[start, end)`` range of row offsets within a table."""
26
+
27
+ start: int
28
+ end: int
29
+
30
+ @property
31
+ def num_rows(self) -> int:
32
+ return self.end - self.start
33
+
34
+
35
+ def plan_offset_shards(num_rows: int, parallelism: int) -> list[OffsetRange]:
36
+ """Split ``[0, num_rows)`` into at most ``parallelism`` contiguous shards.
37
+
38
+ Rows are distributed as evenly as possible: the first ``num_rows %
39
+ parallelism`` shards receive one extra row. Empty shards are never
40
+ produced, so the result is empty when ``num_rows`` is zero and has
41
+ ``num_rows`` single-row entries when ``parallelism`` exceeds ``num_rows``.
42
+
43
+ Args:
44
+ num_rows: Total number of rows in the table. Must not be negative.
45
+ parallelism: Desired number of shards. Must be positive.
46
+
47
+ Returns:
48
+ Contiguous, non-overlapping ranges covering ``[0, num_rows)`` in order.
49
+ """
50
+ if num_rows < 0:
51
+ raise ValueError(f"num_rows must not be negative, got {num_rows}")
52
+ if parallelism <= 0:
53
+ raise ValueError(f"parallelism must be positive, got {parallelism}")
54
+
55
+ if num_rows == 0:
56
+ return []
57
+
58
+ num_shards = min(parallelism, num_rows)
59
+ base, remainder = divmod(num_rows, num_shards)
60
+
61
+ shards: list[OffsetRange] = []
62
+ start = 0
63
+ for i in range(num_shards):
64
+ size = base + (1 if i < remainder else 0)
65
+ shards.append(OffsetRange(start, start + size))
66
+ start += size
67
+
68
+ # The loop above is exact arithmetic; this guards against future edits.
69
+ assert start == num_rows, f"shards covered {start} of {num_rows} rows"
70
+ return shards
71
+
72
+
73
+ def chunk_offsets(offsets: OffsetRange, batch_size: int) -> Iterator[list[int]]:
74
+ """Materialise an offset range as batches of explicit row offsets.
75
+
76
+ ``Table.take_offsets`` requires an explicit list of offsets rather than a
77
+ range. For a large shard that list can be enormous, so it is generated
78
+ lazily inside the worker -- only the two integers of the ``OffsetRange``
79
+ ever cross the wire -- and yielded in ``batch_size`` chunks so no single
80
+ request carries the whole shard.
81
+
82
+ Args:
83
+ offsets: The half-open range to expand.
84
+ batch_size: Maximum number of offsets per yielded chunk. Must be positive.
85
+
86
+ Yields:
87
+ Lists of row offsets, each of length ``batch_size`` except possibly the last.
88
+ """
89
+ if batch_size <= 0:
90
+ raise ValueError(f"batch_size must be positive, got {batch_size}")
91
+
92
+ for start in range(offsets.start, offsets.end, batch_size):
93
+ end = min(start + batch_size, offsets.end)
94
+ yield list(range(start, end))
95
+
96
+
97
+ def split_arrow_table(num_rows: int, max_rows: int) -> list[OffsetRange]:
98
+ """Split ``num_rows`` into consecutive slices of at most ``max_rows``.
99
+
100
+ Used by the write path to break an accumulated batch into request-sized
101
+ pieces before handing them to LanceDB.
102
+ """
103
+ if num_rows < 0:
104
+ raise ValueError(f"num_rows must not be negative, got {num_rows}")
105
+ if max_rows <= 0:
106
+ raise ValueError(f"max_rows must be positive, got {max_rows}")
107
+
108
+ return [
109
+ OffsetRange(start, min(start + max_rows, num_rows))
110
+ for start in range(0, num_rows, max_rows)
111
+ ]
112
+
113
+
114
+ def rows_within_byte_budget(num_rows: int, nbytes: int, max_bytes: int) -> int:
115
+ """Largest slice of a batch whose share of ``nbytes`` fits in ``max_bytes``.
116
+
117
+ A row count does not bound memory. One row of a 1536-dimension float32
118
+ embedding is 6KB and one row carrying a JPEG is unbounded, so a schema wide
119
+ enough exhausts a worker long before any row limit is reached. Sizing a
120
+ slice by the batch's own average row width is what makes the ceiling
121
+ actually about bytes.
122
+
123
+ The estimate is uniform across rows, so a batch with a few outlying rows can
124
+ still overshoot -- this is a ceiling on the common case, not a hard bound.
125
+
126
+ Args:
127
+ num_rows: Rows in the batch being split.
128
+ nbytes: In-memory size of that batch.
129
+ max_bytes: Byte budget for one slice. Must be positive.
130
+
131
+ Returns:
132
+ A row count in ``[1, num_rows]``, or ``0`` when the batch is empty.
133
+ Never zero for a non-empty batch: a single row that exceeds the budget
134
+ on its own still has to be written, and refusing it would deadlock the
135
+ write rather than bound it.
136
+ """
137
+ if num_rows < 0:
138
+ raise ValueError(f"num_rows must not be negative, got {num_rows}")
139
+ if nbytes < 0:
140
+ raise ValueError(f"nbytes must not be negative, got {nbytes}")
141
+ if max_bytes <= 0:
142
+ raise ValueError(f"max_bytes must be positive, got {max_bytes}")
143
+
144
+ if num_rows == 0 or nbytes <= max_bytes:
145
+ return num_rows
146
+
147
+ # Integer arithmetic throughout: at these magnitudes a float ratio can
148
+ # round a slice up past the budget it was supposed to enforce.
149
+ return max(1, (max_bytes * num_rows) // nbytes)
lancedb_ray/_retry.py ADDED
@@ -0,0 +1,196 @@
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ """Retry helpers for transient LanceDB failures.
3
+
4
+ Remote (Cloud/Enterprise) calls go over HTTP and fail transiently; local writes
5
+ contend on the dataset commit lock. Both are worth retrying with backoff, and
6
+ neither should be retried forever.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import logging
12
+ import random
13
+ import re
14
+ import time
15
+ from collections.abc import Callable
16
+
17
+ logger = logging.getLogger(__name__)
18
+
19
+ __all__ = ["RetryPolicy", "call_with_retry", "is_commit_conflict", "is_transient"]
20
+
21
+ # Substrings identifying errors that are worth another attempt. Matched against
22
+ # the lowercased ``str()`` of the exception because the LanceDB Python SDK
23
+ # surfaces server-side failures as generic exception types with descriptive
24
+ # messages rather than a dedicated exception hierarchy.
25
+ _TRANSIENT_MARKERS = (
26
+ "timeout",
27
+ "timed out",
28
+ "connection reset",
29
+ "connection aborted",
30
+ "connection refused",
31
+ "broken pipe",
32
+ "temporarily unavailable",
33
+ "service unavailable",
34
+ "too many requests",
35
+ "internal server error",
36
+ "bad gateway",
37
+ "gateway timeout",
38
+ )
39
+
40
+
41
+ def _status_pattern(*codes: str) -> re.Pattern[str]:
42
+ """Match an HTTP status code as a whole number, not as digits anywhere.
43
+
44
+ A bare substring test reads ``429`` out of "dimension 1429" and ``503`` out
45
+ of "row 8503", classifying a deterministic schema error as retryable. Word
46
+ boundaries keep a code from matching inside a longer number, which is how
47
+ row counts, dimensions and byte sizes reach these messages.
48
+ """
49
+ return re.compile(rf"\b(?:{'|'.join(codes)})\b")
50
+
51
+
52
+ _TRANSIENT_STATUS = _status_pattern("429", "502", "503", "504")
53
+
54
+ _COMMIT_CONFLICT_MARKERS = (
55
+ "commit conflict",
56
+ "concurrent",
57
+ "version already exists",
58
+ "retryable commit",
59
+ "commit was rejected",
60
+ )
61
+
62
+
63
+ def is_transient(error: BaseException) -> bool:
64
+ """Return whether ``error`` looks like a retryable transient failure."""
65
+ message = str(error).lower()
66
+ if any(marker in message for marker in _TRANSIENT_MARKERS):
67
+ return True
68
+ return _TRANSIENT_STATUS.search(message) is not None
69
+
70
+
71
+ #: Failures that mean the request never reached the service, or was rejected
72
+ #: without being applied. Re-sending these cannot duplicate anything.
73
+ _NOT_APPLIED_MARKERS = (
74
+ "connection refused",
75
+ "temporarily unavailable",
76
+ "service unavailable",
77
+ "too many requests",
78
+ "name or service not known",
79
+ "nodename nor servname",
80
+ "failed to resolve",
81
+ )
82
+
83
+ _NOT_APPLIED_STATUS = _status_pattern("429", "503")
84
+
85
+
86
+ def is_definitely_not_applied(error: BaseException) -> bool:
87
+ """Whether ``error`` proves the write never took effect.
88
+
89
+ A read timeout or a dropped connection is ambiguous: the service may have
90
+ committed and only the response was lost. Re-sending a non-idempotent
91
+ append in that case silently duplicates rows, so appends retry only on
92
+ this narrower class.
93
+ """
94
+ message = str(error).lower()
95
+ if any(marker in message for marker in _NOT_APPLIED_MARKERS):
96
+ return True
97
+ return _NOT_APPLIED_STATUS.search(message) is not None
98
+
99
+
100
+ #: Arrow-rs refuses to import a buffer whose pointer is not aligned for its
101
+ #: scalar type. Ray hands out zero-copy views into its object store, and a
102
+ #: view can land unaligned for a type that needs more than 8 bytes, such as
103
+ #: decimal128. The import fails before anything is written.
104
+ _ALIGNMENT_MARKERS = ("is not aligned with the specified scalar type",)
105
+
106
+
107
+ def is_arrow_alignment_error(error: BaseException) -> bool:
108
+ """Whether ``error`` is arrow-rs rejecting an unaligned FFI buffer.
109
+
110
+ Recoverable by copying the batch into freshly allocated memory, and safe
111
+ to retry: the import fails before any data is committed.
112
+ """
113
+ message = str(error)
114
+ return any(marker in message for marker in _ALIGNMENT_MARKERS)
115
+
116
+
117
+ def is_commit_conflict(error: BaseException) -> bool:
118
+ """Return whether ``error`` looks like a losing race on a dataset commit."""
119
+ message = str(error).lower()
120
+ return any(marker in message for marker in _COMMIT_CONFLICT_MARKERS)
121
+
122
+
123
+ class RetryPolicy:
124
+ """Exponential backoff with full jitter.
125
+
126
+ Args:
127
+ max_attempts: Total attempts including the first. ``1`` disables retrying.
128
+ initial_backoff_s: Backoff before the second attempt.
129
+ max_backoff_s: Ceiling on the backoff between attempts.
130
+ predicate: Returns whether a given exception should be retried.
131
+ """
132
+
133
+ def __init__(
134
+ self,
135
+ *,
136
+ max_attempts: int = 5,
137
+ initial_backoff_s: float = 0.5,
138
+ max_backoff_s: float = 32.0,
139
+ predicate: Callable[[BaseException], bool] = is_transient,
140
+ ) -> None:
141
+ if max_attempts < 1:
142
+ raise ValueError(f"max_attempts must be at least 1, got {max_attempts}")
143
+ self.max_attempts = max_attempts
144
+ self.initial_backoff_s = initial_backoff_s
145
+ self.max_backoff_s = max_backoff_s
146
+ self.predicate = predicate
147
+
148
+ def backoff_for(self, attempt: int) -> float:
149
+ """Return the sleep, in seconds, before attempt number ``attempt`` (1-based)."""
150
+ uncapped = self.initial_backoff_s * (2 ** (attempt - 1))
151
+ return random.uniform(0.0, min(uncapped, self.max_backoff_s))
152
+
153
+
154
+ def call_with_retry[T](
155
+ fn: Callable[[], T],
156
+ policy: RetryPolicy,
157
+ *,
158
+ description: str,
159
+ sleep: Callable[[float], None] = time.sleep,
160
+ ) -> T:
161
+ """Call ``fn``, retrying while ``policy.predicate`` accepts the raised error.
162
+
163
+ Args:
164
+ fn: Zero-argument callable to invoke.
165
+ policy: Attempt count, backoff schedule and retry predicate.
166
+ description: Human-readable operation name used in log messages.
167
+ sleep: Injectable sleep, so tests do not spend real time backing off.
168
+
169
+ Returns:
170
+ Whatever ``fn`` returns on its first successful attempt.
171
+
172
+ Raises:
173
+ BaseException: The final error, once attempts are exhausted or the
174
+ predicate rejects it.
175
+ """
176
+ last_error: BaseException | None = None
177
+ for attempt in range(1, policy.max_attempts + 1):
178
+ try:
179
+ return fn()
180
+ except Exception as error: # noqa: BLE001 - re-raised below
181
+ last_error = error
182
+ if attempt == policy.max_attempts or not policy.predicate(error):
183
+ raise
184
+ backoff = policy.backoff_for(attempt)
185
+ logger.warning(
186
+ "%s failed (attempt %d/%d): %s. Retrying in %.2fs.",
187
+ description,
188
+ attempt,
189
+ policy.max_attempts,
190
+ error,
191
+ backoff,
192
+ )
193
+ sleep(backoff)
194
+
195
+ # Unreachable: the loop either returns or raises.
196
+ raise AssertionError(f"retry loop exited without result: {last_error}")