fix(python): coordinate native table state transitions

This commit is contained in:
Gatefixer
2026-08-06 10:11:29 +00:00
parent 62dea8acd8
commit 6a1dacb1fe
2 changed files with 252 additions and 44 deletions
+106 -43
View File
@@ -7,6 +7,7 @@ import asyncio
import inspect
import deprecation
import os
import threading
import warnings
from abc import ABC, abstractmethod
from dataclasses import dataclass
@@ -2190,6 +2191,7 @@ class LanceTable(Table):
"""Capture the state needed to replace inherited native handles."""
self._name = name
self._pid = os.getpid()
self._native_state_guard = (self._pid, threading.RLock())
# A native table owns object-store clients and connection pools. Those
# handles must not be used after fork, so retain a process-independent
@@ -2259,6 +2261,21 @@ class LanceTable(Table):
else:
self._legacy_checkout_version = value
def _native_state_lock(self):
"""Return the per-process lock coordinating native mode and reopen state."""
pid = os.getpid()
guard = getattr(self, "_native_state_guard", None)
if guard is None:
candidate = (pid, threading.RLock())
guard = self.__dict__.setdefault("_native_state_guard", candidate)
elif guard[0] != pid:
# A lock inherited while another parent thread held it cannot be
# safely acquired in the child. Child state starts single-threaded,
# so replace it before coordinating the first reopen.
guard = (pid, threading.RLock())
self._native_state_guard = guard
return guard[1]
@classmethod
def _open_from_reopen_state(
cls,
@@ -2307,40 +2324,41 @@ class LanceTable(Table):
def _ensure_open(self) -> None:
"""Reopen native table handles inherited from another process."""
pid = os.getpid()
if getattr(self, "_pid", pid) == pid:
return
with self._native_state_lock():
pid = os.getpid()
if getattr(self, "_pid", pid) == pid:
return
state = getattr(self, "_reopen_state", None)
if (
state is None
or not state.can_reopen_after_fork
or state.connection_state is None
):
# In-memory and opaque Rust-only connections cannot be recreated
# from connection metadata. Their local handles retain the prior
# best-effort fork behavior.
state = getattr(self, "_reopen_state", None)
if (
state is None
or not state.can_reopen_after_fork
or state.connection_state is None
):
# In-memory and opaque Rust-only connections cannot be recreated
# from connection metadata. Their local handles retain the prior
# best-effort fork behavior.
self._pid = pid
return
from lancedb import deserialize_conn
connection = deserialize_conn(state.connection_state, for_worker=True)
reopened = self._open_from_reopen_state(
connection,
state,
)
# Keep this Python object stable because user datasets commonly retain
# it across fork. Replace every process-bound component with the fresh
# child's equivalent.
self._conn = reopened._conn
self._table = reopened._table
self._namespace_client = reopened._namespace_client
self._pushdown_operations = reopened._pushdown_operations
self._route_pushdown_to_rust = reopened._route_pushdown_to_rust
self._reopen_state = reopened._reopen_state
self._pid = pid
return
from lancedb import deserialize_conn
connection = deserialize_conn(state.connection_state, for_worker=True)
reopened = self._open_from_reopen_state(
connection,
state,
)
# Keep this Python object stable because user datasets commonly retain
# it across fork. Replace every process-bound component with the fresh
# child's equivalent.
self._conn = reopened._conn
self._table = reopened._table
self._namespace_client = reopened._namespace_client
self._pushdown_operations = reopened._pushdown_operations
self._route_pushdown_to_rust = reopened._route_pushdown_to_rust
self._reopen_state = reopened._reopen_state
self._pid = pid
@property
def name(self) -> str:
@@ -2592,6 +2610,48 @@ class LanceTable(Table):
raise ValueError(str(err)) from err
raise
async def _commit_native_state(
self,
transition,
version: Optional[int],
started: threading.Event,
finished: threading.Event,
):
"""Commit a native transition and its fork coordinate as one task."""
started.set()
try:
task = asyncio.ensure_future(transition)
try:
result = await asyncio.shield(task)
except asyncio.CancelledError:
# BackgroundEventLoop cancels its submitted task when the
# waiting caller is interrupted. Let an already-started native
# transition reach its authoritative terminal state before the
# per-table boundary is released.
result = await task
self._checkout_version = version
raise
self._checkout_version = version
return result
finally:
finished.set()
def _run_native_state_transition(self, transition, version: Optional[int]):
started = threading.Event()
finished = threading.Event()
try:
return LOOP.run(
self._commit_native_state(transition, version, started, finished)
)
except BaseException:
if started.is_set():
while not finished.is_set():
try:
finished.wait()
except BaseException: # noqa: PERF203
continue
raise
def checkout(self, version: Union[int, str]):
"""Checkout a version of the table. This is an in-place operation.
@@ -2631,9 +2691,11 @@ class LanceTable(Table):
# 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
with self._native_state_lock():
resolved_version = self._resolve_checkout_version(version)
self._run_native_state_transition(
self._table.checkout(resolved_version), resolved_version
)
def checkout_latest(self):
"""Checkout the latest version of the table. This is an in-place operation.
@@ -2641,8 +2703,8 @@ class LanceTable(Table):
The table will be set back into standard mode, and will track the latest
version of the table.
"""
LOOP.run(self._table.checkout_latest())
self._checkout_version = None
with self._native_state_lock():
self._run_native_state_transition(self._table.checkout_latest(), None)
def restore(self, version: Optional[Union[int, str]] = None):
"""Restore a version of the table. This is an in-place operation.
@@ -2688,12 +2750,13 @@ class LanceTable(Table):
>>> len(table.list_versions())
4
"""
if version is not None:
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
with self._native_state_lock():
if version is not None:
resolved_version = self._resolve_checkout_version(version)
self._run_native_state_transition(
self._table.checkout(resolved_version), resolved_version
)
self._run_native_state_transition(self._table.restore(), None)
def count_rows(self, filter: Optional[str] = None) -> int:
return LOOP.run(self._table.count_rows(filter))
+146 -1
View File
@@ -2,6 +2,7 @@
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import asyncio
import ctypes
import gc
import os
@@ -9,7 +10,7 @@ import sys
import threading
import warnings
import weakref
from concurrent.futures import ThreadPoolExecutor
from concurrent.futures import CancelledError, ThreadPoolExecutor
from datetime import date, datetime, timedelta
from time import sleep
from typing import List
@@ -2238,6 +2239,150 @@ def test_string_tag_resolution_failure_does_not_mutate_handle(operation):
assert inner.restore_calls == 0
def test_native_state_transitions_are_serialized(monkeypatch):
from lancedb.background_loop import LOOP
class Inner:
def __init__(self):
self.live_version = None
async def checkout(self, version):
self.live_version = version
inner = Inner()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = None
first_native_done = threading.Event()
release_first_call = threading.Event()
second_call_started = threading.Event()
second_call_done = threading.Event()
errors = []
original_run = LOOP.run
def delayed_delivery(awaitable):
result = original_run(awaitable)
if threading.current_thread().name == "checkout-1":
first_native_done.set()
assert release_first_call.wait(5)
return result
monkeypatch.setattr(LOOP, "run", delayed_delivery)
def checkout(version):
if version == 2:
second_call_started.set()
try:
table.checkout(version)
except BaseException as err:
errors.append(err)
finally:
if version == 2:
second_call_done.set()
first = threading.Thread(target=checkout, args=(1,), name="checkout-1")
first.start()
assert first_native_done.wait(5)
second = threading.Thread(target=checkout, args=(2,), name="checkout-2")
second.start()
assert second_call_started.wait(5)
assert not second_call_done.wait(0.1)
release_first_call.set()
first.join(5)
second.join(5)
assert not first.is_alive()
assert not second.is_alive()
assert errors == []
assert inner.live_version == 2
assert table._checkout_version == 2
@pytest.mark.parametrize(
("operation", "args", "initial_version", "expected_version"),
[
("checkout", (11,), 3, 11),
("checkout_latest", (), 3, None),
("restore", (), 11, None),
],
)
def test_native_state_commits_before_success_delivery(
monkeypatch, operation, args, initial_version, expected_version
):
from lancedb.background_loop import LOOP
class Inner:
def __init__(self, live_version):
self.live_version = live_version
async def checkout(self, version):
self.live_version = version
async def checkout_latest(self):
self.live_version = None
async def restore(self):
self.live_version = None
inner = Inner(initial_version)
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = initial_version
original_run = LOOP.run
def success_then_interrupt(awaitable):
original_run(awaitable)
raise KeyboardInterrupt("injected after native success")
monkeypatch.setattr(LOOP, "run", success_then_interrupt)
with pytest.raises(KeyboardInterrupt, match="injected after native success"):
getattr(table, operation)(*args)
assert inner.live_version == expected_version
assert table._checkout_version == expected_version
def test_native_state_waits_for_cancelled_delivery(monkeypatch):
from lancedb.background_loop import LOOP
class Inner:
def __init__(self):
self.live_version = 3
async def checkout(self, version):
await asyncio.sleep(0.01)
self.live_version = version
inner = Inner()
table = LanceTable.__new__(LanceTable)
table._table = inner
table._checkout_version = 3
original_run = LOOP.run
def cancel_while_running(awaitable):
async def cancel_after_start():
task = asyncio.create_task(awaitable)
await asyncio.sleep(0)
task.cancel()
return await task
return original_run(cancel_after_start())
monkeypatch.setattr(LOOP, "run", cancel_while_running)
with pytest.raises(CancelledError):
table.checkout(11)
assert inner.live_version == 11
assert table._checkout_version == 11
def test_reopen_preserves_explicit_table_location(tmp_path):
db = lancedb.connect(tmp_path / "db")
location = str(tmp_path / "physical-table")