mirror of
https://github.com/lancedb/lancedb.git
synced 2026-06-29 00:50:38 +00:00
Compare commits
1 Commits
jack/node-
...
jack/pytho
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
956a8ee714 |
@@ -89,6 +89,8 @@ def connect(
|
|||||||
If presented, connect to LanceDB cloud.
|
If presented, connect to LanceDB cloud.
|
||||||
Otherwise, connect to a database on file system or cloud storage.
|
Otherwise, connect to a database on file system or cloud storage.
|
||||||
Can be set via environment variable `LANCEDB_API_KEY`.
|
Can be set via environment variable `LANCEDB_API_KEY`.
|
||||||
|
OAuth configuration is currently supported only by ``connect_async``;
|
||||||
|
synchronous LanceDB Cloud connections require an API key.
|
||||||
region: str, default "us-east-1"
|
region: str, default "us-east-1"
|
||||||
The region to use for LanceDB Cloud.
|
The region to use for LanceDB Cloud.
|
||||||
host_override: str, optional
|
host_override: str, optional
|
||||||
@@ -340,6 +342,7 @@ async def connect_async(
|
|||||||
session: Optional[Session] = None,
|
session: Optional[Session] = None,
|
||||||
manifest_enabled: bool = False,
|
manifest_enabled: bool = False,
|
||||||
namespace_client_properties: Optional[Dict[str, str]] = None,
|
namespace_client_properties: Optional[Dict[str, str]] = None,
|
||||||
|
oauth_config=None,
|
||||||
) -> AsyncConnection:
|
) -> AsyncConnection:
|
||||||
"""Connect to a LanceDB database.
|
"""Connect to a LanceDB database.
|
||||||
|
|
||||||
@@ -389,6 +392,10 @@ async def connect_async(
|
|||||||
namespace_client_properties : dict, optional
|
namespace_client_properties : dict, optional
|
||||||
Additional directory namespace client properties to use with
|
Additional directory namespace client properties to use with
|
||||||
``manifest_enabled=True``.
|
``manifest_enabled=True``.
|
||||||
|
oauth_config : OAuthConfig, optional
|
||||||
|
OAuth configuration for LanceDB Cloud/Enterprise. This is supported by
|
||||||
|
``connect_async`` only; synchronous ``connect`` uses API key
|
||||||
|
authentication for ``db://`` URIs.
|
||||||
|
|
||||||
Examples
|
Examples
|
||||||
--------
|
--------
|
||||||
@@ -435,6 +442,7 @@ async def connect_async(
|
|||||||
session,
|
session,
|
||||||
manifest_enabled,
|
manifest_enabled,
|
||||||
namespace_client_properties,
|
namespace_client_properties,
|
||||||
|
oauth_config,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -280,6 +280,7 @@ async def connect(
|
|||||||
session: Optional[Session],
|
session: Optional[Session],
|
||||||
manifest_enabled: bool = False,
|
manifest_enabled: bool = False,
|
||||||
namespace_client_properties: Optional[Dict[str, str]] = None,
|
namespace_client_properties: Optional[Dict[str, str]] = None,
|
||||||
|
oauth_config: Optional[Any] = None,
|
||||||
) -> Connection: ...
|
) -> Connection: ...
|
||||||
|
|
||||||
class RecordBatchStream:
|
class RecordBatchStream:
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ from typing import List, Optional
|
|||||||
from lancedb import __version__
|
from lancedb import __version__
|
||||||
|
|
||||||
from .header import HeaderProvider
|
from .header import HeaderProvider
|
||||||
|
from .oauth import OAuthConfig, OAuthFlowType
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"TimeoutConfig",
|
"TimeoutConfig",
|
||||||
@@ -16,6 +17,8 @@ __all__ = [
|
|||||||
"TlsConfig",
|
"TlsConfig",
|
||||||
"ClientConfig",
|
"ClientConfig",
|
||||||
"HeaderProvider",
|
"HeaderProvider",
|
||||||
|
"OAuthConfig",
|
||||||
|
"OAuthFlowType",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
75
python/python/lancedb/remote/oauth.py
Normal file
75
python/python/lancedb/remote/oauth.py
Normal file
@@ -0,0 +1,75 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from enum import Enum
|
||||||
|
from typing import List, Optional
|
||||||
|
|
||||||
|
|
||||||
|
class OAuthFlowType(str, Enum):
|
||||||
|
"""OAuth authentication flow types."""
|
||||||
|
|
||||||
|
CLIENT_CREDENTIALS = "client_credentials"
|
||||||
|
"""Client Credentials grant (service-to-service / M2M)."""
|
||||||
|
|
||||||
|
AZURE_MANAGED_IDENTITY = "azure_managed_identity"
|
||||||
|
"""Azure Managed Identity via IMDS."""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class OAuthConfig:
|
||||||
|
"""OAuth configuration for LanceDB authentication.
|
||||||
|
|
||||||
|
All token acquisition and refresh is handled in the Rust layer.
|
||||||
|
This config is passed through to Rust via PyO3.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
issuer_url : str
|
||||||
|
OIDC issuer URL or OAuth authority URL.
|
||||||
|
For Azure: ``https://login.microsoftonline.com/{tenant_id}/v2.0``
|
||||||
|
client_id : str
|
||||||
|
Application / Client ID.
|
||||||
|
scopes : List[str]
|
||||||
|
OAuth scopes to request.
|
||||||
|
For Azure managed identity, exactly one scope or resource is required.
|
||||||
|
For example: ``["api://{app_id}/.default"]``
|
||||||
|
flow : OAuthFlowType
|
||||||
|
Authentication flow to use. Default: CLIENT_CREDENTIALS.
|
||||||
|
client_secret : Optional[str]
|
||||||
|
Client secret (required for CLIENT_CREDENTIALS).
|
||||||
|
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.
|
||||||
|
|
||||||
|
Examples
|
||||||
|
--------
|
||||||
|
Client Credentials (service-to-service):
|
||||||
|
|
||||||
|
>>> config = OAuthConfig(
|
||||||
|
... issuer_url="https://login.microsoftonline.com/{tenant}/v2.0",
|
||||||
|
... client_id="app-id",
|
||||||
|
... client_secret="secret",
|
||||||
|
... scopes=["api://lancedb-api/.default"],
|
||||||
|
... )
|
||||||
|
|
||||||
|
Azure Managed Identity:
|
||||||
|
|
||||||
|
>>> config = OAuthConfig(
|
||||||
|
... issuer_url="https://login.microsoftonline.com/{tenant}/v2.0",
|
||||||
|
... client_id="app-id",
|
||||||
|
... scopes=["api://lancedb-api/.default"],
|
||||||
|
... flow=OAuthFlowType.AZURE_MANAGED_IDENTITY,
|
||||||
|
... )
|
||||||
|
"""
|
||||||
|
|
||||||
|
issuer_url: str
|
||||||
|
client_id: str
|
||||||
|
scopes: List[str]
|
||||||
|
flow: OAuthFlowType = OAuthFlowType.CLIENT_CREDENTIALS
|
||||||
|
client_secret: Optional[str] = field(default=None, repr=False)
|
||||||
|
managed_identity_client_id: Optional[str] = None
|
||||||
|
refresh_buffer_secs: Optional[int] = None
|
||||||
@@ -539,7 +539,7 @@ impl Connection {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[pyfunction]
|
#[pyfunction]
|
||||||
#[pyo3(signature = (uri, api_key=None, region=None, host_override=None, read_consistency_interval=None, client_config=None, storage_options=None, session=None, manifest_enabled=false, namespace_client_properties=None))]
|
#[pyo3(signature = (uri, api_key=None, region=None, host_override=None, read_consistency_interval=None, client_config=None, storage_options=None, session=None, manifest_enabled=false, namespace_client_properties=None, oauth_config=None))]
|
||||||
#[allow(clippy::too_many_arguments)]
|
#[allow(clippy::too_many_arguments)]
|
||||||
pub fn connect(
|
pub fn connect(
|
||||||
py: Python<'_>,
|
py: Python<'_>,
|
||||||
@@ -553,6 +553,7 @@ pub fn connect(
|
|||||||
session: Option<crate::session::Session>,
|
session: Option<crate::session::Session>,
|
||||||
manifest_enabled: bool,
|
manifest_enabled: bool,
|
||||||
namespace_client_properties: Option<HashMap<String, String>>,
|
namespace_client_properties: Option<HashMap<String, String>>,
|
||||||
|
oauth_config: Option<crate::oauth::PyOAuthConfig>,
|
||||||
) -> PyResult<Bound<'_, PyAny>> {
|
) -> PyResult<Bound<'_, PyAny>> {
|
||||||
future_into_py(py, async move {
|
future_into_py(py, async move {
|
||||||
let mut builder = lancedb::connect(&uri);
|
let mut builder = lancedb::connect(&uri);
|
||||||
@@ -582,6 +583,11 @@ pub fn connect(
|
|||||||
if let Some(client_config) = client_config {
|
if let Some(client_config) = client_config {
|
||||||
builder = builder.client_config(client_config.into());
|
builder = builder.client_config(client_config.into());
|
||||||
}
|
}
|
||||||
|
if let Some(oauth_config) = oauth_config {
|
||||||
|
let config: lancedb::remote::oauth::OAuthConfig =
|
||||||
|
oauth_config.try_into().infer_error()?;
|
||||||
|
builder = builder.oauth_config(config);
|
||||||
|
}
|
||||||
if let Some(session) = session {
|
if let Some(session) = session {
|
||||||
builder = builder.session(session.inner.clone());
|
builder = builder.session(session.inner.clone());
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ pub mod expr;
|
|||||||
pub mod header;
|
pub mod header;
|
||||||
pub mod index;
|
pub mod index;
|
||||||
pub mod namespace;
|
pub mod namespace;
|
||||||
|
pub mod oauth;
|
||||||
pub mod permutation;
|
pub mod permutation;
|
||||||
pub mod query;
|
pub mod query;
|
||||||
pub mod runtime;
|
pub mod runtime;
|
||||||
|
|||||||
72
python/src/oauth.rs
Normal file
72
python/src/oauth.rs
Normal file
@@ -0,0 +1,72 @@
|
|||||||
|
// SPDX-License-Identifier: Apache-2.0
|
||||||
|
// SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||||
|
|
||||||
|
use pyo3::FromPyObject;
|
||||||
|
|
||||||
|
use lancedb::error::Error;
|
||||||
|
use lancedb::remote::oauth::{OAuthConfig, OAuthFlow};
|
||||||
|
|
||||||
|
/// Python-side OAuth configuration, extracted via FromPyObject.
|
||||||
|
/// Maps to `lancedb.remote.oauth.OAuthConfig` Python dataclass.
|
||||||
|
#[derive(FromPyObject)]
|
||||||
|
pub struct PyOAuthConfig {
|
||||||
|
pub issuer_url: String,
|
||||||
|
pub client_id: String,
|
||||||
|
pub scopes: Vec<String>,
|
||||||
|
pub flow: String,
|
||||||
|
pub client_secret: Option<String>,
|
||||||
|
pub managed_identity_client_id: Option<String>,
|
||||||
|
pub refresh_buffer_secs: Option<u64>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl TryFrom<PyOAuthConfig> for OAuthConfig {
|
||||||
|
type Error = Error;
|
||||||
|
|
||||||
|
fn try_from(py: PyOAuthConfig) -> Result<Self, Self::Error> {
|
||||||
|
let flow = match py.flow.as_str() {
|
||||||
|
"client_credentials" => OAuthFlow::ClientCredentials,
|
||||||
|
"azure_managed_identity" => OAuthFlow::AzureManagedIdentity {
|
||||||
|
client_id: py.managed_identity_client_id,
|
||||||
|
},
|
||||||
|
other => {
|
||||||
|
return Err(Error::InvalidInput {
|
||||||
|
message: format!("Unknown OAuth flow type: {other}"),
|
||||||
|
});
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
Ok(Self {
|
||||||
|
issuer_url: py.issuer_url,
|
||||||
|
client_id: py.client_id,
|
||||||
|
client_secret: py.client_secret,
|
||||||
|
scopes: py.scopes,
|
||||||
|
flow,
|
||||||
|
refresh_buffer_secs: py.refresh_buffer_secs,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn test_unknown_oauth_flow_returns_invalid_input() {
|
||||||
|
let config = PyOAuthConfig {
|
||||||
|
issuer_url: "https://issuer.example.com".to_string(),
|
||||||
|
client_id: "client-id".to_string(),
|
||||||
|
scopes: vec!["scope".to_string()],
|
||||||
|
flow: "typo".to_string(),
|
||||||
|
client_secret: None,
|
||||||
|
managed_identity_client_id: None,
|
||||||
|
refresh_buffer_secs: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let err = OAuthConfig::try_from(config).unwrap_err();
|
||||||
|
assert!(matches!(
|
||||||
|
err,
|
||||||
|
Error::InvalidInput { message }
|
||||||
|
if message == "Unknown OAuth flow type: typo"
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
33
python/tests/test_oauth.py
Normal file
33
python/tests/test_oauth.py
Normal file
@@ -0,0 +1,33 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
|
||||||
|
|
||||||
|
import importlib.util
|
||||||
|
import sys
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
def _load_oauth_module():
|
||||||
|
oauth_path = (
|
||||||
|
Path(__file__).parents[1] / "python" / "lancedb" / "remote" / "oauth.py"
|
||||||
|
)
|
||||||
|
spec = importlib.util.spec_from_file_location("lancedb_remote_oauth", oauth_path)
|
||||||
|
module = importlib.util.module_from_spec(spec)
|
||||||
|
assert spec.loader is not None
|
||||||
|
sys.modules[spec.name] = module
|
||||||
|
spec.loader.exec_module(module)
|
||||||
|
return module
|
||||||
|
|
||||||
|
|
||||||
|
def test_oauth_config_repr_redacts_client_secret():
|
||||||
|
oauth = _load_oauth_module()
|
||||||
|
|
||||||
|
config = oauth.OAuthConfig(
|
||||||
|
issuer_url="https://issuer.example.com",
|
||||||
|
client_id="client-id",
|
||||||
|
scopes=["scope"],
|
||||||
|
client_secret="super-secret",
|
||||||
|
)
|
||||||
|
|
||||||
|
rendered = repr(config)
|
||||||
|
assert "super-secret" not in rendered
|
||||||
|
assert "client_secret" not in rendered
|
||||||
Reference in New Issue
Block a user