From 6a1dacb1fea3edd1d7de15a9254ed433204fd825 Mon Sep 17 00:00:00 2001 From: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 10:11:29 +0000 Subject: [PATCH] fix(python): coordinate native table state transitions --- python/python/lancedb/table.py | 149 +++++++++++++++++++++--------- python/python/tests/test_table.py | 147 ++++++++++++++++++++++++++++- 2 files changed, 252 insertions(+), 44 deletions(-) diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index 635d694b9..14d086fdf 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -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)) diff --git a/python/python/tests/test_table.py b/python/python/tests/test_table.py index fbb116783..fdd0a8aa7 100644 --- a/python/python/tests/test_table.py +++ b/python/python/tests/test_table.py @@ -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")