feat: add persistent OAuth token cache and session APIs (#4182)

Stacked on #4173 (diff includes it until that merges; will rebase
after). Addresses the token-cache part of [Colin's
review](https://github.com/lancedb/lancedb/pull/4173#issuecomment-5674048100).

Adds an explicit, opt-in persistent OAuth token cache shared by Rust,
Python, and Node clients, plus `login` / `status` / `logout` session
APIs, so short-lived processes (CLIs, scripts, notebooks) reuse one
session instead of restarting a browser or device flow on every start.

- **Opt-in and minimal**: existing callers stay memory-only and lazy.
Only refresh tokens are persisted (never access tokens, never client
secrets), so there are no local token-expiry decisions to get wrong when
clocks move. Each process start performs one silent refresh grant.
- **Hardened file backend**: private directory (`0700`), per-record
files (`0600`), owner validation, symlink rejection, and atomic `rename`
replacement. Corrupt, truncated, unknown-version, or permission-invalid
records fail with actionable errors naming the file. Native keyring
backends were evaluated (keyring crate routes Linux through D-Bus/zbus:
heavy deps, headless/CI flakiness) and are deferred; the file store is
the explicit opt-in, not a downgrade from a keyring.
- **Cache key**: SHA-256 of the canonical identity (issuer, client ID,
sorted/de-duplicated scopes, flow, public/confidential), so no secret
appears in a filename and distinct identities never collide. Versioned
record schema (`version: 1`). One record per identity: last login wins,
documented.
- **Cross-process rotation locking**: per-key `fs4` file lock (`flock` /
`LockFileEx`) around the refresh critical section — acquire, reread the
durable record, refresh exactly once, atomically store the rotated
refresh token, release. The OS releases locks on process death, so
crashes cannot strand stale locks. Only confirmed
`invalid_grant`/`invalid_token` deletes a record and reauthenticates;
transport, 5xx, 429, and parse failures retain it.
- **Session APIs**: `OAuthSession::login/status/logout` in Rust,
`lancedb.remote.OAuthSession` (async) in Python, `OAuthSession` class in
Node. `status` returns non-secret metadata only. `logout` removes only
the local credential — provider revocation (RFC 7009) is a deliberate
follow-up, and local logout never terminates browser SSO. Azure managed
identity is rejected for persistence (machine identity stays in memory);
client credentials have nothing refreshable to persist and stay
memory-only.
- No CLI binary exists in this repo, so this ships library APIs plus doc
examples in all three languages.

Tests: Rust unit + mock-IdP integration (cache-key
canonicalization/separation, record
versioning/corruption/truncation/symlink/owner/perms, lock serialization
+ release, two concurrent providers proving no `invalid_grant` and
correct rotation, transient-failure retention, `invalid_grant` delete +
reauthenticate, login/status/logout lifecycle, client-credentials no-op,
IMDS rejection, secret redaction); Python lifecycle + a true
two-subprocess cross-process reuse test (second process refreshes once,
never hits the device endpoint); Node lifecycle + device-flow login
test. Local builds were skipped in development; CI validates all
bindings.

---------

Co-authored-by: Xuanwo <github@xuanwo.io>
This commit is contained in:
Jack Ye
2026-09-16 02:01:10 +08:00
committed by GitHub
co-authored by Xuanwo
parent 575286922b
commit 2f88b71c21
25 changed files with 5464 additions and 97 deletions
+27
View File
@@ -267,6 +267,33 @@ class JobInfo:
@property
def created_at_millis(self) -> int: ...
class SessionStatus:
@property
def refreshable(self) -> bool: ...
@property
def issuer_url(self) -> str: ...
@property
def client_id(self) -> str: ...
@property
def scopes(self) -> List[str]: ...
@property
def flow(self) -> str: ...
@property
def obtained_at(self) -> Optional[int]: ...
def __repr__(self) -> str: ...
class SessionLogout:
@property
def removed(self) -> bool: ...
def __repr__(self) -> str: ...
class OAuthSession:
def __init__(self, config: Any) -> None: ...
async def login(self) -> SessionStatus: ...
async def status(self) -> SessionStatus: ...
async def logout(self) -> SessionLogout: ...
def __repr__(self) -> str: ...
class JobFailureInfo:
@property
def phase(self) -> Optional[str]: ...
+3 -1
View File
@@ -9,7 +9,7 @@ from typing import List, Optional
from lancedb import __version__
from .header import HeaderProvider
from .oauth import OAuthConfig, OAuthFlowType
from .oauth import OAuthConfig, OAuthFlowType, OAuthSession, TokenCacheOptions
# The API reference renders this module with a single mkdocstrings directive,
# which only picks up names listed here. New public names must be added to this
@@ -22,6 +22,8 @@ __all__ = [
"HeaderProvider",
"OAuthConfig",
"OAuthFlowType",
"OAuthSession",
"TokenCacheOptions",
]
+138
View File
@@ -12,10 +12,49 @@ class OAuthFlowType(str, Enum):
CLIENT_CREDENTIALS = "client_credentials"
"""Client Credentials grant (service-to-service / M2M)."""
AUTHORIZATION_CODE = "authorization_code"
"""Interactive Authorization Code grant, using PKCE by default."""
DEVICE_CODE = "device_code"
"""Device Authorization grant for CLI and headless environments."""
AZURE_MANAGED_IDENTITY = "azure_managed_identity"
"""Azure Managed Identity via IMDS."""
@dataclass
class TokenCacheOptions:
"""Options for the persistent OAuth token cache.
The cache is opt-in: it is only used when set as ``token_cache`` on
:class:`OAuthConfig`. Only refresh tokens are persisted, in a private
directory with owner-only permissions, so short-lived processes can reuse
an authenticated session instead of re-prompting on every start.
Parameters
----------
cache_dir : Optional[str]
Directory that holds cached credentials. Defaults to
``$XDG_CACHE_HOME/lancedb/oauth``, ``$HOME/.cache/lancedb/oauth`` on
Unix, or ``%LOCALAPPDATA%\\lancedb\\oauth`` on Windows. The directory
is created with owner-only permissions (``0700``) when missing.
lock_timeout_secs : Optional[int]
How long to wait for the cross-process refresh lock before failing
(default: 30 seconds).
Examples
--------
>>> opts = TokenCacheOptions(cache_dir="/tmp/my-app/oauth-cache")
Multiple identities (issuer, client, scopes, flow, client
authentication) get separate cache entries. Within one identity the most
recent login wins.
"""
cache_dir: Optional[str] = None
lock_timeout_secs: Optional[int] = None
@dataclass
class OAuthConfig:
"""OAuth configuration for LanceDB authentication.
@@ -38,12 +77,23 @@ class OAuthConfig:
Authentication flow to use. Default: CLIENT_CREDENTIALS.
client_secret : Optional[str]
Client secret (required for CLIENT_CREDENTIALS).
redirect_uri : Optional[str]
Loopback redirect URI for AUTHORIZATION_CODE. The default is
``http://127.0.0.1:{callback_port}/callback``.
callback_port : Optional[int]
Port for the AUTHORIZATION_CODE loopback callback server (default: 8400).
use_pkce : bool
Protect AUTHORIZATION_CODE with S256 PKCE (default: True).
managed_identity_client_id : Optional[str]
Client ID for user-assigned managed identity (AZURE_MANAGED_IDENTITY).
refresh_buffer_secs : Optional[int]
Seconds before expiry to trigger proactive refresh (default: 300).
Keep this well below the token TTL; if it is greater than or equal to
the TTL, each request refreshes the token.
token_cache : Optional[TokenCacheOptions]
Opt in to the persistent token cache so short-lived processes reuse
one session. Only supported by AUTHORIZATION_CODE and DEVICE_CODE;
azure managed identity is rejected. Default: None (memory only).
Examples
--------
@@ -64,6 +114,29 @@ class OAuthConfig:
... scopes=["api://lancedb-api/.default"],
... flow=OAuthFlowType.AZURE_MANAGED_IDENTITY,
... )
Authorization Code with PKCE:
The authorization URL is written to standard error before LanceDB tries to
open a browser, so it can be copied in headless environments.
>>> config = OAuthConfig(
... issuer_url="https://login.microsoftonline.com/{tenant}/v2.0",
... client_id="app-id",
... scopes=["openid", "api://lancedb-api/access"],
... flow=OAuthFlowType.AUTHORIZATION_CODE,
... )
Device Authorization with a persistent cache, so later processes reuse
the session without a new device prompt:
>>> config = OAuthConfig(
... issuer_url="https://login.microsoftonline.com/{tenant}/v2.0",
... client_id="app-id",
... scopes=["openid", "offline_access", "api://lancedb-api/access"],
... flow=OAuthFlowType.DEVICE_CODE,
... token_cache=TokenCacheOptions(),
... )
"""
issuer_url: str
@@ -71,5 +144,70 @@ class OAuthConfig:
scopes: List[str]
flow: OAuthFlowType = OAuthFlowType.CLIENT_CREDENTIALS
client_secret: Optional[str] = field(default=None, repr=False)
redirect_uri: Optional[str] = None
callback_port: Optional[int] = None
use_pkce: bool = True
managed_identity_client_id: Optional[str] = None
refresh_buffer_secs: Optional[int] = None
token_cache: Optional[TokenCacheOptions] = None
class OAuthSession:
"""Explicit OAuth session lifecycle for the persistent token cache.
Built from the same :class:`OAuthConfig` used for
:func:`lancedb.connect_async` (including its ``token_cache`` options).
A connection created with the same configuration shares the cache, so
logging in here prepares tokens for later processes without any database
request.
``login`` always runs the configured interactive flow and replaces the
cached session (the most recent login wins). ``logout`` removes only the
local credential; it does not revoke anything with the provider and does
not sign out of a browser SSO session.
Examples
--------
>>> config = OAuthConfig(
... issuer_url="https://issuer.example.com",
... client_id="my-app",
... scopes=["openid", "offline_access"],
... flow=OAuthFlowType.DEVICE_CODE,
... token_cache=TokenCacheOptions(),
... )
>>> session = OAuthSession(config) # doctest: +SKIP
>>> status = await session.login() # doctest: +SKIP
>>> status.refreshable # doctest: +SKIP
True
"""
def __init__(self, config: OAuthConfig):
from lancedb._lancedb import OAuthSession as PyOAuthSession
self._inner: PyOAuthSession = PyOAuthSession(config)
async def login(self):
"""Eagerly run the configured flow and store the session.
Returns a :class:`lancedb._lancedb.SessionStatus` describing the
cached session. A successful login always replaces any prior cached
session for this identity; if the provider does not issue a refresh
token (for example without ``offline_access``), the previous record
is removed and ``refreshable`` is ``False``.
"""
return await self._inner.login()
async def status(self):
"""Report whether a cached session exists, with safe metadata.
Never contacts the identity provider and never exposes token values.
"""
return await self._inner.status()
async def logout(self):
"""Remove the matching local cached credential.
Returns a :class:`lancedb._lancedb.SessionLogout` whose ``removed``
flag reports whether a credential existed. Logout is idempotent.
"""
return await self._inner.logout()
+3
View File
@@ -47,6 +47,9 @@ pub fn _lancedb(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_class::<Connection>()?;
m.add_class::<Session>()?;
m.add_class::<Table>()?;
m.add_class::<crate::oauth::PyOAuthSession>()?;
m.add_class::<crate::oauth::PySessionStatus>()?;
m.add_class::<crate::oauth::PySessionLogout>()?;
m.add_class::<crate::job::Job>()?;
m.add_class::<crate::job::JobInfo>()?;
m.add_class::<crate::job::JobDescription>()?;
+246 -6
View File
@@ -1,10 +1,33 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
use pyo3::FromPyObject;
use std::path::PathBuf;
use std::sync::Arc;
use pyo3::{FromPyObject, PyResult, Python, pyclass, pymethods};
use crate::error::PythonErrorExt;
use crate::runtime::future_into_py;
use lancedb::error::Error;
use lancedb::remote::oauth::{OAuthConfig, OAuthFlow};
use lancedb::remote::oauth::{AuthorizationCodeOptions, OAuthConfig, OAuthFlow};
use lancedb::remote::{OAuthSession, SessionLogout, SessionStatus, TokenCacheOptions};
/// Python-side persistent token cache options, extracted via FromPyObject.
/// Maps to `lancedb.remote.oauth.TokenCacheOptions` Python dataclass.
#[derive(FromPyObject, Default)]
pub struct PyTokenCacheOptions {
pub cache_dir: Option<String>,
pub lock_timeout_secs: Option<u64>,
}
impl From<PyTokenCacheOptions> for TokenCacheOptions {
fn from(py: PyTokenCacheOptions) -> Self {
TokenCacheOptions {
cache_dir: py.cache_dir.map(PathBuf::from),
lock_timeout_secs: py.lock_timeout_secs,
}
}
}
/// Python-side OAuth configuration, extracted via FromPyObject.
/// Maps to `lancedb.remote.oauth.OAuthConfig` Python dataclass.
@@ -15,8 +38,12 @@ pub struct PyOAuthConfig {
pub scopes: Vec<String>,
pub flow: String,
pub client_secret: Option<String>,
pub redirect_uri: Option<String>,
pub callback_port: Option<u16>,
pub use_pkce: bool,
pub managed_identity_client_id: Option<String>,
pub refresh_buffer_secs: Option<u64>,
pub token_cache: Option<PyTokenCacheOptions>,
}
impl TryFrom<PyOAuthConfig> for OAuthConfig {
@@ -25,6 +52,17 @@ impl TryFrom<PyOAuthConfig> for OAuthConfig {
fn try_from(py: PyOAuthConfig) -> Result<Self, Self::Error> {
let flow = match py.flow.as_str() {
"client_credentials" => OAuthFlow::ClientCredentials,
"authorization_code" => {
let mut options = AuthorizationCodeOptions::new().use_pkce(py.use_pkce);
if let Some(redirect_uri) = py.redirect_uri {
options = options.redirect_uri(redirect_uri);
}
if let Some(callback_port) = py.callback_port {
options = options.callback_port(callback_port);
}
OAuthFlow::AuthorizationCode(options)
}
"device_code" => OAuthFlow::DeviceCode,
"azure_managed_identity" => OAuthFlow::AzureManagedIdentity {
client_id: py.managed_identity_client_id,
},
@@ -42,6 +80,148 @@ impl TryFrom<PyOAuthConfig> for OAuthConfig {
scopes: py.scopes,
flow,
refresh_buffer_secs: py.refresh_buffer_secs,
token_cache: py.token_cache.map(TokenCacheOptions::from),
})
}
}
/// Wrapper around [`lancedb::remote::SessionStatus`] exposing safe metadata.
#[pyclass(name = "SessionStatus", skip_from_py_object)]
#[derive(Clone)]
pub struct PySessionStatus {
inner: SessionStatus,
}
#[pymethods]
impl PySessionStatus {
/// Whether a cached session exists that can obtain tokens without
/// interactive authentication.
#[getter]
pub fn refreshable(&self) -> bool {
self.inner.refreshable
}
/// Canonical issuer URL of the cached session.
#[getter]
pub fn issuer_url(&self) -> String {
self.inner.issuer_url.clone()
}
/// Client ID of the cached session.
#[getter]
pub fn client_id(&self) -> String {
self.inner.client_id.clone()
}
/// Canonical (sorted, de-duplicated) scopes of the cached session.
#[getter]
pub fn scopes(&self) -> Vec<String> {
self.inner.scopes.clone()
}
/// Flow that produced the cached session.
#[getter]
pub fn flow(&self) -> String {
self.inner.flow.clone()
}
/// When the cached session was obtained, as Unix seconds.
#[getter]
pub fn obtained_at(&self) -> Option<u64> {
self.inner.obtained_at
}
pub fn __repr__(&self) -> String {
format!(
"SessionStatus(refreshable={}, issuer_url='{}', client_id='{}', flow='{}')",
self.inner.refreshable, self.inner.issuer_url, self.inner.client_id, self.inner.flow
)
}
}
impl From<SessionStatus> for PySessionStatus {
fn from(inner: SessionStatus) -> Self {
Self { inner }
}
}
/// Wrapper around [`lancedb::remote::SessionLogout`].
#[pyclass(name = "SessionLogout", skip_from_py_object)]
#[derive(Clone)]
pub struct PySessionLogout {
inner: SessionLogout,
}
#[pymethods]
impl PySessionLogout {
/// Whether a cached credential was removed.
#[getter]
pub fn removed(&self) -> bool {
self.inner.removed
}
pub fn __repr__(&self) -> String {
format!("SessionLogout(removed={})", self.inner.removed)
}
}
impl From<SessionLogout> for PySessionLogout {
fn from(inner: SessionLogout) -> Self {
Self { inner }
}
}
/// Wrapper around [`lancedb::remote::OAuthSession`].
#[pyclass(name = "OAuthSession", skip_from_py_object)]
#[derive(Clone)]
pub struct PyOAuthSession {
inner: Arc<OAuthSession>,
}
#[pymethods]
impl PyOAuthSession {
/// Create a session manager for the given OAuth configuration.
///
/// The configuration must set ``token_cache`` options and use a flow that
/// supports persistent sessions (authorization code or device code).
#[new]
pub fn new(config: PyOAuthConfig) -> PyResult<Self> {
let config: OAuthConfig = config.try_into().infer_error()?;
let inner = OAuthSession::new(config).infer_error()?;
Ok(Self {
inner: Arc::new(inner),
})
}
/// Eagerly run the configured authentication flow and store the session.
pub fn login<'py>(&self, py: Python<'py>) -> PyResult<pyo3::Bound<'py, pyo3::PyAny>> {
let inner = Arc::clone(&self.inner);
future_into_py(py, async move {
inner.login().await.map(PySessionStatus::from).infer_error()
})
}
/// Report whether a matching cached session exists, with safe metadata.
pub fn status<'py>(&self, py: Python<'py>) -> PyResult<pyo3::Bound<'py, pyo3::PyAny>> {
let inner = Arc::clone(&self.inner);
future_into_py(py, async move {
inner
.status()
.await
.map(PySessionStatus::from)
.infer_error()
})
}
/// Remove the matching local cached credential.
pub fn logout<'py>(&self, py: Python<'py>) -> PyResult<pyo3::Bound<'py, pyo3::PyAny>> {
let inner = Arc::clone(&self.inner);
future_into_py(py, async move {
inner
.logout()
.await
.map(PySessionLogout::from)
.infer_error()
})
}
}
@@ -50,16 +230,27 @@ impl TryFrom<PyOAuthConfig> for OAuthConfig {
mod tests {
use super::*;
#[test]
fn test_unknown_oauth_flow_returns_invalid_input() {
let config = PyOAuthConfig {
fn base_config() -> PyOAuthConfig {
PyOAuthConfig {
issuer_url: "https://issuer.example.com".to_string(),
client_id: "client-id".to_string(),
scopes: vec!["scope".to_string()],
flow: "typo".to_string(),
flow: "device_code".to_string(),
client_secret: None,
redirect_uri: None,
callback_port: None,
use_pkce: true,
managed_identity_client_id: None,
refresh_buffer_secs: None,
token_cache: None,
}
}
#[test]
fn test_unknown_oauth_flow_returns_invalid_input() {
let config = PyOAuthConfig {
flow: "typo".to_string(),
..base_config()
};
let err = OAuthConfig::try_from(config).unwrap_err();
@@ -69,4 +260,53 @@ mod tests {
if message == "Unknown OAuth flow type: typo"
));
}
#[test]
fn test_authorization_code_conversion_preserves_options() {
let config = PyOAuthConfig {
flow: "authorization_code".to_string(),
client_secret: Some("secret".to_string()),
redirect_uri: Some("http://127.0.0.1:9000/callback".to_string()),
callback_port: Some(9000),
use_pkce: false,
..base_config()
};
let converted = OAuthConfig::try_from(config).unwrap();
let OAuthFlow::AuthorizationCode(options) = converted.flow else {
panic!("expected authorization code flow");
};
assert_eq!(
options.redirect_uri.as_deref(),
Some("http://127.0.0.1:9000/callback")
);
assert_eq!(options.callback_port, Some(9000));
assert!(!options.use_pkce);
}
#[test]
fn test_device_code_conversion() {
let config = base_config();
let converted = OAuthConfig::try_from(config).unwrap();
assert!(matches!(converted.flow, OAuthFlow::DeviceCode));
}
#[test]
fn test_token_cache_conversion() {
let config = PyOAuthConfig {
token_cache: Some(PyTokenCacheOptions {
cache_dir: Some("/tmp/oauth-cache".to_string()),
lock_timeout_secs: Some(5),
}),
..base_config()
};
let converted = OAuthConfig::try_from(config).unwrap();
let cache = converted.token_cache.expect("token cache options");
assert_eq!(
cache.cache_dir.as_deref(),
Some(std::path::Path::new("/tmp/oauth-cache"))
);
assert_eq!(cache.lock_timeout_secs, Some(5));
}
}
+295
View File
@@ -1,10 +1,19 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import asyncio
import importlib.util
import json
import os
import subprocess
import sys
import threading
import urllib.parse
from http.server import BaseHTTPRequestHandler, HTTPServer
from pathlib import Path
import pytest
def _load_oauth_module():
oauth_path = (
@@ -31,3 +40,289 @@ def test_oauth_config_repr_redacts_client_secret():
rendered = repr(config)
assert "super-secret" not in rendered
assert "client_secret" not in rendered
def test_authorization_code_uses_pkce_by_default():
oauth = _load_oauth_module()
config = oauth.OAuthConfig(
issuer_url="https://issuer.example.com",
client_id="client-id",
scopes=["openid"],
flow=oauth.OAuthFlowType.AUTHORIZATION_CODE,
)
assert config.use_pkce is True
assert config.redirect_uri is None
assert config.callback_port is None
def test_device_code_flow_value():
oauth = _load_oauth_module()
assert oauth.OAuthFlowType.DEVICE_CODE.value == "device_code"
def test_token_cache_options_default_to_memory_only():
oauth = _load_oauth_module()
config = oauth.OAuthConfig(
issuer_url="https://issuer.example.com",
client_id="client-id",
scopes=["openid"],
)
assert config.token_cache is None
options = oauth.TokenCacheOptions()
assert options.cache_dir is None
assert options.lock_timeout_secs is None
def _remote_oauth():
pytest.importorskip("lancedb")
from lancedb.remote import oauth as remote_oauth
return remote_oauth
def _device_config(remote_oauth, issuer_url, cache_dir):
return remote_oauth.OAuthConfig(
issuer_url=issuer_url,
client_id="client-id",
scopes=["openid"],
flow=remote_oauth.OAuthFlowType.DEVICE_CODE,
token_cache=remote_oauth.TokenCacheOptions(cache_dir=str(cache_dir)),
)
def test_oauth_session_status_and_logout_without_cache_entry(tmp_path):
remote_oauth = _remote_oauth()
config = _device_config(remote_oauth, "https://issuer.example.com", tmp_path)
session = remote_oauth.OAuthSession(config)
status = asyncio.run(session.status())
assert status.refreshable is False
assert status.issuer_url == "https://issuer.example.com"
assert status.client_id == "client-id"
assert status.scopes == ["openid"]
assert status.flow == "device_code"
assert status.obtained_at is None
logout = asyncio.run(session.logout())
assert logout.removed is False
class _MockIdpState:
def __init__(self, port):
self.port = port
self.lock = threading.Lock()
self.device_authorizations = 0
self.refresh_grants = 0
self.invalid_grant_rejections = 0
self.access_tokens_issued = 0
self.current_refresh = None
class _MockIdpHandler(BaseHTTPRequestHandler):
@property
def state(self) -> _MockIdpState:
return self.server.state
def log_message(self, fmt, *args):
pass
def _respond(self, status, payload):
body = json.dumps(payload).encode()
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def do_GET(self):
if self.path.startswith("/.well-known/openid-configuration"):
base = f"http://127.0.0.1:{self.state.port}"
self._respond(
200,
{
"token_endpoint": f"{base}/token",
"device_authorization_endpoint": f"{base}/device",
},
)
else:
self._respond(404, {})
def do_POST(self):
length = int(self.headers.get("Content-Length", 0))
body = self.rfile.read(length).decode()
params = urllib.parse.parse_qs(body)
if self.path == "/device":
with self.state.lock:
self.state.device_authorizations += 1
base = f"http://127.0.0.1:{self.state.port}"
self._respond(
200,
{
"device_code": "device-code",
"user_code": "ABCD-EFGH",
"verification_uri": f"{base}/verify",
"expires_in": 60,
"interval": 1,
},
)
return
if self.path == "/token":
grant_type = params.get("grant_type", [""])[0]
with self.state.lock:
if grant_type == "refresh_token":
self.state.refresh_grants += 1
offered = params.get("refresh_token", [""])[0]
if offered != self.state.current_refresh:
self.state.invalid_grant_rejections += 1
self._respond(400, {"error": "invalid_grant"})
return
elif "device_code" not in grant_type:
self._respond(400, {"error": "unsupported_grant_type"})
return
self.state.access_tokens_issued += 1
number = self.state.access_tokens_issued
refresh = f"refresh-{number}"
self.state.current_refresh = refresh
self._respond(
200,
{
"access_token": f"access-{number}",
"refresh_token": refresh,
"expires_in": 3600,
},
)
return
self._respond(404, {})
def _start_mock_idp() -> tuple[_MockIdpState, HTTPServer]:
server = HTTPServer(("127.0.0.1", 0), _MockIdpHandler)
state = _MockIdpState(server.server_address[1])
server.state = state
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
return state, server
def _run_subprocess(script: Path, issuer_url: str, cache_dir: Path):
env = dict(os.environ)
env["LANCEDB_OAUTH_BROWSER"] = "/usr/bin/true"
result = subprocess.run(
[sys.executable, str(script), issuer_url, str(cache_dir)],
capture_output=True,
text=True,
timeout=120,
env=env,
)
assert result.returncode == 0, (
f"subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"
)
return result
LOGIN_SCRIPT = """
import asyncio
import sys
from lancedb.remote import OAuthConfig, OAuthFlowType, OAuthSession, TokenCacheOptions
issuer_url, cache_dir = sys.argv[1], sys.argv[2]
config = OAuthConfig(
issuer_url=issuer_url,
client_id="client-id",
scopes=["openid"],
flow=OAuthFlowType.DEVICE_CODE,
token_cache=TokenCacheOptions(cache_dir=cache_dir),
)
session = OAuthSession(config)
status = asyncio.run(session.login())
assert status.refreshable, "login must cache a refresh token"
print("LOGIN-OK")
"""
REUSE_SCRIPT = """
import asyncio
import sys
import lancedb
from lancedb.remote import OAuthConfig, OAuthFlowType, OAuthSession, TokenCacheOptions
issuer_url, cache_dir = sys.argv[1], sys.argv[2]
config = OAuthConfig(
issuer_url=issuer_url,
client_id="client-id",
scopes=["openid"],
flow=OAuthFlowType.DEVICE_CODE,
token_cache=TokenCacheOptions(cache_dir=cache_dir),
)
session = OAuthSession(config)
status = asyncio.run(session.status())
assert status.refreshable, "second process must see the cached session"
async def main():
# Point the database endpoint at a dead port. OAuth headers are fetched
# before the request is sent, so a successful refresh proves the second
# process reused the cached session; only the database call fails.
db = await lancedb.connect_async(
"db://e2e",
host_override="http://127.0.0.1:1",
client_config={"retry_config": {"retries": 0}},
oauth_config=config,
)
try:
await db.table_names()
except Exception:
print("DATABASE-UNREACHABLE-AS-EXPECTED")
else:
raise AssertionError("expected the database request to fail")
asyncio.run(main())
print("REUSE-OK")
"""
def test_cross_process_session_reuse_without_new_prompt(tmp_path):
pytest.importorskip("lancedb")
state, server = _start_mock_idp()
try:
issuer_url = f"http://127.0.0.1:{state.port}"
login_script = tmp_path / "login.py"
login_script.write_text(LOGIN_SCRIPT)
reuse_script = tmp_path / "reuse.py"
reuse_script.write_text(REUSE_SCRIPT)
cache_dir = tmp_path / "oauth-cache"
result = _run_subprocess(login_script, issuer_url, cache_dir)
assert "LOGIN-OK" in result.stdout
assert state.device_authorizations == 1
result = _run_subprocess(reuse_script, issuer_url, cache_dir)
assert "REUSE-OK" in result.stdout
assert "DATABASE-UNREACHABLE-AS-EXPECTED" in result.stdout
# The second process refreshed exactly once and never started a new
# interactive device flow.
assert state.refresh_grants == 1
assert state.device_authorizations == 1
assert state.invalid_grant_rejections == 0
logout = asyncio.run(
_remote_oauth()
.OAuthSession(_device_config(_remote_oauth(), issuer_url, cache_dir))
.logout()
)
assert logout.removed is True
finally:
server.shutdown()
server.server_close()