mirror of
https://github.com/lancedb/lancedb.git
synced 2026-09-10 23:32:35 +00:00
## Summary
Add SQL execution to remote LanceDB connections. On the standard
synchronous connection, `execute_query` waits for the initial result
stream and returns its Arrow reader. `execute_query_async` is called
without Python `await` and immediately returns a query handle for status
inspection, streaming, or cancellation. Local databases report that SQL
is not supported.
The transport and query lifecycle live in Rust. Python exposes
native-backed synchronous and asynchronous connection methods and query
wrappers; it does not use PyArrow's Flight client.
## User experience
The standard synchronous connection supports both direct reads and
background query execution:
```python
db = lancedb.connect(
"db://analytics",
api_key="ldb_...",
sql_host_override="grpc+tls://sql.example.com:10026",
)
# Direct execution waits only until the initial result stream is available.
# Later batches continue streaming as the query progresses.
reader = db.execute_query(
"SELECT * FROM events",
default_namespace_path=["production"],
)
for batch in reader:
print(batch.num_rows)
# Background execution returns a query handle immediately. Despite the
# `_async` suffix, no Python `await` is needed on a synchronous connection.
query = db.execute_query_async("SELECT * FROM events")
print(query.id)
description = db.describe_query(query.id)
print(description.status)
print(description.progress)
print(description.expires_at)
# Start reading as soon as the service advertises partial results. The reader
# continues polling and yields newly available record batches until the query
# and all result endpoints are complete.
reader = query.reader()
for batch in reader:
print(batch.num_rows)
# Or cancel a different still-running query. Its status becomes "cancelling"
# while the server is still working, then "cancelled" once confirmed.
cancelled_query = db.execute_query_async("SELECT * FROM large_events")
cancelled_query.cancel()
```
The less commonly used asynchronous connection exposes the same
operations as coroutines:
```python
async_db = await lancedb.connect_async(
"db://analytics",
api_key="ldb_...",
sql_host_override="grpc+tls://sql.example.com:10026",
)
query = await async_db.execute_query_async("SELECT * FROM events")
async for batch in await query.reader():
print(batch.num_rows)
```
The UUIDv7 query id is scoped to the connection that submitted it. The
connection retains lightweight shared query state used by
`query.describe()` and `db.describe_query(query.id)`; the id does not
encode SQL or a Flight continuation token and is not a cross-connection
resume token. Abandoned state has bounded retention, and terminal state
remains available briefly.
Unqualified table names use the connected database and the `public`
namespace by default. `default_namespace_path` accepts a list such as
`["production", "events"]`. SQL can still use qualified names to
reference other databases and namespaces available to the deployment.
## Design
- Uses Arrow Flight `PollFlightInfo` for submission and long polling,
`DoGet` for results, and `CancelFlightInfo` for cancellation. Each
`PollInfo.info` is treated as the cumulative set of currently available
endpoints, so advertised tickets are consumed once and batches can be
delivered before execution is complete.
- Serializes result completion and cancellation into one lifecycle. A
server-accepted request reports `cancelling` and wakes blocked
status/result work; a later retry can confirm `cancelled`. Result
retrieval is rejected after cancellation is accepted, while cancellation
after a result was already delivered is a no-op.
- Assigns a time-ordered UUIDv7 connection-scoped query id and retains
only shared evolving lifecycle state, keeping SQL, Flight continuation
tokens, and Arrow result data out of public ids and the registry.
- Leaves admission control to the server while honoring server
expiration and a local fallback retention window for abandoned entries.
- Retains terminal ids for five minutes so they remain available for
connection-level description.
- Keeps one lazily initialized SQL client on each remote database
connection and attaches fresh authentication, routing, namespace, and
request metadata to every operation.
- Applies the configured overall timeout to each execution, description,
reader, and cancellation operation. A result reader carries one absolute
deadline from `reader()` through the end of streaming; connect and read
timeouts continue to bound their individual phases.
- Returns a bounded, backpressured, single-consumer Arrow stream rather
than collecting the full result in memory. Dropping the reader stops
downloading but does not implicitly cancel the server query.
- Preserves typed schemas for empty result sets through the stream
schema.
- Accepts Flight result messages up to 1 GiB so a valid row containing a
large blob, string, or vector is not rejected by tonic's 4 MiB default
receive limit.
- Supports the Python client first while keeping the authoritative
implementation in the Rust core.
184 lines
6.3 KiB
Python
184 lines
6.3 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
|
|
|
"""Header providers for LanceDB remote connections.
|
|
|
|
This module provides a flexible header management framework for LanceDB remote
|
|
connections, allowing users to implement custom header strategies for
|
|
authentication, request tracking, custom metadata, or any other header-based
|
|
requirements.
|
|
|
|
The module includes the HeaderProvider abstract base class and example implementations
|
|
(StaticHeaderProvider and OAuthProvider) that demonstrate common patterns.
|
|
|
|
The HeaderProvider interface is designed to be called before each request to the remote
|
|
server, enabling dynamic header scenarios where values may need to be
|
|
refreshed, rotated, or computed on-demand.
|
|
"""
|
|
|
|
from abc import ABC, abstractmethod
|
|
from typing import Dict, Optional, Callable, Any
|
|
import time
|
|
import threading
|
|
|
|
|
|
class HeaderProvider(ABC):
|
|
"""Abstract base class for providing custom headers for each request.
|
|
|
|
Users can implement this interface to provide dynamic headers for various purposes
|
|
such as authentication (OAuth tokens, API keys), request tracking (correlation IDs),
|
|
custom metadata, or any other header-based requirements. The provider is called
|
|
before each request to ensure fresh header values are always used.
|
|
|
|
Error Handling
|
|
--------------
|
|
If get_headers() raises an exception, the request will fail. Implementations
|
|
should handle recoverable errors internally (e.g., retry token refresh) and
|
|
only raise exceptions for unrecoverable errors.
|
|
"""
|
|
|
|
@abstractmethod
|
|
def get_headers(self) -> Dict[str, str]:
|
|
"""Get the latest headers to be added to requests.
|
|
|
|
This method is called before each request to the remote LanceDB server.
|
|
Implementations should return headers that will be merged with existing headers.
|
|
|
|
Returns
|
|
-------
|
|
Dict[str, str]
|
|
Dictionary of header names to values to add to the request.
|
|
|
|
Raises
|
|
------
|
|
Exception
|
|
If unable to fetch headers, the exception will be propagated
|
|
and the request will fail.
|
|
"""
|
|
pass
|
|
|
|
|
|
class StaticHeaderProvider(HeaderProvider):
|
|
"""Example implementation: A simple header provider that returns static headers.
|
|
|
|
This is an example implementation showing how to create a HeaderProvider
|
|
for cases where headers don't change during the session. Users can use this
|
|
as a reference for implementing their own providers.
|
|
|
|
Parameters
|
|
----------
|
|
headers : Dict[str, str]
|
|
Static headers to return for every request.
|
|
"""
|
|
|
|
def __init__(self, headers: Dict[str, str]):
|
|
"""Initialize with static headers.
|
|
|
|
Parameters
|
|
----------
|
|
headers : Dict[str, str]
|
|
Headers to return for every request.
|
|
"""
|
|
self._headers = headers.copy()
|
|
|
|
def get_headers(self) -> Dict[str, str]:
|
|
"""Return the static headers.
|
|
|
|
Returns
|
|
-------
|
|
Dict[str, str]
|
|
Copy of the static headers.
|
|
"""
|
|
return self._headers.copy()
|
|
|
|
|
|
class OAuthProvider(HeaderProvider):
|
|
"""Example implementation: OAuth token provider with automatic refresh.
|
|
|
|
This is an example implementation showing how to manage OAuth tokens
|
|
with automatic refresh when they expire. Users can use this as a reference
|
|
for implementing their own OAuth or token-based authentication providers.
|
|
|
|
Parameters
|
|
----------
|
|
token_fetcher : Callable[[], Dict[str, Any]]
|
|
Function that fetches a new token. Should return a dict with
|
|
'access_token' and optionally 'expires_in' (seconds until expiration).
|
|
refresh_buffer_seconds : int, optional
|
|
Number of seconds before expiration to trigger refresh. Default is 300
|
|
(5 minutes).
|
|
"""
|
|
|
|
def __init__(
|
|
self, token_fetcher: Callable[[], Any], refresh_buffer_seconds: int = 300
|
|
):
|
|
"""Initialize the OAuth provider.
|
|
|
|
Parameters
|
|
----------
|
|
token_fetcher : Callable[[], Any]
|
|
Function to fetch new tokens. Should return dict with
|
|
'access_token' and optionally 'expires_in'.
|
|
refresh_buffer_seconds : int, optional
|
|
Seconds before expiry to refresh token. Default 300.
|
|
"""
|
|
self._token_fetcher = token_fetcher
|
|
self._refresh_buffer = refresh_buffer_seconds
|
|
self._current_token: Optional[str] = None
|
|
self._token_expires_at: Optional[float] = None
|
|
self._refresh_lock = threading.Lock()
|
|
|
|
def _refresh_token_if_needed(self) -> None:
|
|
"""Refresh the token if it's expired or close to expiring."""
|
|
with self._refresh_lock:
|
|
# Check again inside the lock in case another thread refreshed
|
|
if self._needs_refresh():
|
|
token_data = self._token_fetcher()
|
|
|
|
self._current_token = token_data.get("access_token")
|
|
if not self._current_token:
|
|
raise ValueError("Token fetcher did not return 'access_token'")
|
|
|
|
# Set expiration if provided
|
|
expires_in = token_data.get("expires_in")
|
|
if expires_in:
|
|
self._token_expires_at = time.time() + expires_in
|
|
else:
|
|
# Token doesn't expire or expiration unknown
|
|
self._token_expires_at = None
|
|
|
|
def _needs_refresh(self) -> bool:
|
|
"""Check if token needs refresh."""
|
|
if self._current_token is None:
|
|
return True
|
|
|
|
if self._token_expires_at is None:
|
|
# No expiration info, assume token is valid
|
|
return False
|
|
|
|
# Refresh if we're within the buffer time of expiration
|
|
return time.time() >= (self._token_expires_at - self._refresh_buffer)
|
|
|
|
def get_headers(self) -> Dict[str, str]:
|
|
"""Get OAuth headers, refreshing token if needed.
|
|
|
|
Returns
|
|
-------
|
|
Dict[str, str]
|
|
Headers with Bearer token authorization.
|
|
|
|
Raises
|
|
------
|
|
Exception
|
|
If unable to fetch or refresh token.
|
|
"""
|
|
self._refresh_token_if_needed()
|
|
|
|
if not self._current_token:
|
|
raise RuntimeError("Failed to obtain OAuth token")
|
|
|
|
return {
|
|
"Authorization": f"Bearer {self._current_token}",
|
|
"x-lancedb-credential-type": "oidc",
|
|
}
|