fix(python): resolve tags before table checkout

This commit is contained in:
Gatefixer
2026-08-06 09:12:55 +00:00
parent ae9e8e8f8d
commit 62dea8acd8
2 changed files with 92 additions and 8 deletions
+21 -8
View File
@@ -2580,6 +2580,18 @@ class LanceTable(Table):
table._checkout_version = version
return table
def _resolve_checkout_version(self, version: Union[int, str]) -> int:
if isinstance(version, int):
return version
try:
return self.tags.get_version(version)
except RuntimeError as err:
# Native checkout historically exposes an unknown tag as ValueError.
# Preserve that contract while resolving tags before mutating the table.
if "Ref not found" in str(err) and "does not exist" in str(err):
raise ValueError(str(err)) from err
raise
def checkout(self, version: Union[int, str]):
"""Checkout a version of the table. This is an in-place operation.
@@ -2616,10 +2628,12 @@ class LanceTable(Table):
vector type
0 [1.1, 0.9] vector
"""
LOOP.run(self._table.checkout(version))
# Resolve tags to their numeric version so a forked child can reopen
# the same pinned view through ``open_table(version=...)``.
self._checkout_version = version if isinstance(version, int) else self.version
# Resolve tags before mutating the native handle. This leaves the live
# handle and reopen descriptor aligned if tag lookup fails, and avoids a
# second fallible version lookup after checkout succeeds.
resolved_version = self._resolve_checkout_version(version)
LOOP.run(self._table.checkout(resolved_version))
self._checkout_version = resolved_version
def checkout_latest(self):
"""Checkout the latest version of the table. This is an in-place operation.
@@ -2675,10 +2689,9 @@ class LanceTable(Table):
4
"""
if version is not None:
LOOP.run(self._table.checkout(version))
self._checkout_version = (
version if isinstance(version, int) else self.version
)
resolved_version = self._resolve_checkout_version(version)
LOOP.run(self._table.checkout(resolved_version))
self._checkout_version = resolved_version
LOOP.run(self._table.restore())
self._checkout_version = None
+71
View File
@@ -2167,6 +2167,77 @@ def test_restore_tracks_checkout_when_restore_fails():
assert table._checkout_version == inner.live_version
@pytest.mark.parametrize(
("operation", "expected_descriptor", "expected_restore_calls"),
[("checkout", 11, 0), ("restore", None, 1)],
)
def test_string_tag_resolves_before_checkout(
operation, expected_descriptor, expected_restore_calls
):
class Tags:
async def get_version(self, tag):
assert tag == "tag-v1"
return 11
class NoPostCheckoutVersionLookup:
def __init__(self):
self.tags = Tags()
self.checkout_versions = []
self.restore_calls = 0
async def checkout(self, version):
self.checkout_versions.append(version)
async def version(self):
raise RuntimeError("post-checkout version lookup must not run")
async def restore(self):
self.restore_calls += 1
inner = NoPostCheckoutVersionLookup()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = 3
getattr(table, operation)("tag-v1")
assert inner.checkout_versions == [11]
assert table._checkout_version == expected_descriptor
assert inner.restore_calls == expected_restore_calls
@pytest.mark.parametrize("operation", ["checkout", "restore"])
def test_string_tag_resolution_failure_does_not_mutate_handle(operation):
class FailingTags:
async def get_version(self, tag):
assert tag == "missing-tag"
raise RuntimeError("injected tag lookup failure")
class UnchangedTable:
def __init__(self):
self.tags = FailingTags()
self.checkout_calls = 0
self.restore_calls = 0
async def checkout(self, version):
self.checkout_calls += 1
async def restore(self):
self.restore_calls += 1
inner = UnchangedTable()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = 3
with pytest.raises(RuntimeError, match="injected tag lookup failure"):
getattr(table, operation)("missing-tag")
assert table._checkout_version == 3
assert inner.checkout_calls == 0
assert inner.restore_calls == 0
def test_reopen_preserves_explicit_table_location(tmp_path):
db = lancedb.connect(tmp_path / "db")
location = str(tmp_path / "physical-table")