Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/upgrade_dependencies.yml
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ jobs:
matrix:
os: ['ubuntu-latest']
package: ["mp-api"]
python-version: ["3.11", "3.12", "3.13", "3.14"]
python-version: ["3.12", "3.13", "3.14"]
steps:
- uses: actions/checkout@v4
with:
Expand Down
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,6 @@

[![testing](https://github.com/materialsproject/api/actions/workflows/testing.yml/badge.svg?branch=main)](https://github.com/materialsproject/api/actions?query=workflow%3Atesting+branch%3Amain)
[![codecov](https://codecov.io/gh/materialsproject/api/branch/main/graph/badge.svg)](https://codecov.io/gh/materialsproject/api)
![python](https://img.shields.io/badge/Python-3.11+-blue.svg?logo=python&logoColor=white)
![python](https://img.shields.io/badge/Python-3.12+-blue.svg?logo=python&logoColor=white)

This repository is the development environment for the new Materials Project API. A core client implementation will reside here. For information on how to use the API, please see the updated [documentation](https://docs.materialsproject.org/downloading-data/how-do-i-download-the-materials-project-database).
1 change: 0 additions & 1 deletion mp_api/_test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,6 @@ def _normalize(doc, field: str):
if k not in ("_page", "_sort_fields", "chunk_size", "fields")
}
for sort_field in [sort_fields] if isinstance(sort_fields, str) else sort_fields:

asc = search_method(
_page=1,
_sort_fields=sort_field,
Expand Down
3 changes: 2 additions & 1 deletion mp_api/client/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,4 +16,5 @@
except PackageNotFoundError: # pragma: no cover
__version__ = os.getenv("SETUPTOOLS_SCM_PRETEND_VERSION", "")

logging.getLogger(__name__).addHandler(logging.NullHandler())
logger = logging.getLogger(__name__)
logger.addHandler(logging.NullHandler())
184 changes: 115 additions & 69 deletions mp_api/client/core/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
from urllib3.util.retry import Retry

from mp_api.client._server_utils import get_consumer, get_user_api_key, is_dev_env
from mp_api.client.core.delta import DeltaCatalog
from mp_api.client.core.exceptions import (
MPRestError,
MPRestWarning,
Expand All @@ -61,6 +62,8 @@
from collections.abc import Callable, Iterable, Iterator
from typing import Any

from arro3.core import RecordBatchReader

from mp_api.client.core.utils import LazyImport

try:
Expand Down Expand Up @@ -88,13 +91,7 @@
"thermo",
]

hdlr = logging.StreamHandler()
fmt = logging.Formatter("%(name)s - %(levelname)s - %(message)s")
hdlr.setFormatter(fmt)

logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO)
logger.addHandler(hdlr)


def _batched(iterable: Iterable, n: int) -> Iterator:
Expand All @@ -106,28 +103,42 @@ def _batched(iterable: Iterable, n: int) -> Iterator:


class QueryBuilderWithCache(QueryBuilder):
def __init__(self, catalog: DeltaCatalog | None = None, _warn: bool = True) -> None:
"""Deprecated: use `mp_api.client.core.delta.DeltaCatalog`.

def __init__(self) -> None:
"""Extend deltalake.QueryBuilder with stored DeltaTables.

The deltalake.QueryBuilder class does not permit introspection
of registered DeltaTables through the python API.

Re-registering a DeltaTable
(1) wastes time by reading its metadata
(2) raises an exception because a table is already registered
Kept for backwards compatibility. Tables registered here, and
queries run through it, are delegated to a `DeltaCatalog`. Resters
given this object via `query_builder=` share its catalog.

This class simply allows for caching the DeltaTable instances
and table names on the QueryBuilder class.
Args:
catalog (DeltaCatalog or None) : catalog to delegate to.
A new one is created if None.
_warn (bool) : internal, whether to emit a DeprecationWarning
"""
# Dict of table names (labels) to DeltaTable instances
self._delta_tables: dict[str, DeltaTable] = {}
if _warn:
warnings.warn(
"QueryBuilderWithCache is deprecated and will be removed in a future "
"release. Pass `delta_catalog=DeltaCatalog()` "
"(from mp_api.client.core.delta) to MPRester instead.",
category=DeprecationWarning,
stacklevel=2,
)
self.catalog: DeltaCatalog = catalog if catalog is not None else DeltaCatalog()
super().__init__()

@property
def _delta_tables(self) -> dict[str, DeltaTable]:
"""Map of table names (labels) to DeltaTable instances."""
return self.catalog.tables

def register(self, table_name: str, delta_table: DeltaTable) -> QueryBuilder:
"""Register and cache a DeltaTable."""
self._delta_tables[table_name] = delta_table
return super().register(table_name, delta_table)
"""Register a DeltaTable in the underlying catalog."""
self.catalog.add(table_name, delta_table)
return self

def execute(self, sql: str) -> RecordBatchReader:
"""Execute SQL against the tables in the underlying catalog."""
return self.catalog.execute_stream(sql)


class _Rester:
Expand All @@ -148,6 +159,7 @@ def __init__(
) = MAPI_CLIENT_SETTINGS.LOCAL_DATASET_CACHE,
force_renew: bool = False,
query_builder: QueryBuilderWithCache | None = None,
delta_catalog: DeltaCatalog | None = None,
**kwargs,
) -> None:
"""Initialize a RESTer.
Expand Down Expand Up @@ -182,8 +194,12 @@ def __init__(
local_dataset_cache: Target directory for downloading full datasets. Defaults
to 'mp_datasets' in the user's home directory
force_renew: Option to overwrite existing local dataset
query_builder : Instance of QueryBuilderWithCache to use in querying delta tables
query_builder : DEPRECATED, use `delta_catalog`. Instance of QueryBuilderWithCache
whose catalog is used for querying delta tables.
NOTE: Must be a QueryBuilderWithCache, a deltalake.QueryBuilder will be ignored.
delta_catalog : Instance of DeltaCatalog to use for querying delta tables.
Share one instance across resters (e.g. one per web-server worker) to
reuse loaded table snapshots. Takes precedence over `query_builder`.
**kwargs: access to legacy kwargs that may be in the process of being deprecated
"""
self.api_key = get_user_api_key(api_key=api_key)
Expand All @@ -210,6 +226,9 @@ def __init__(
self._query_builder = (
query_builder if isinstance(query_builder, QueryBuilderWithCache) else None
)
if self._query_builder is not None and delta_catalog is None:
delta_catalog = self._query_builder.catalog
self._delta_catalog: DeltaCatalog | None = delta_catalog

if "monty_decode" in kwargs:
# Pop to not repeatedly trigger warning to the user
Expand All @@ -230,9 +249,24 @@ def session(self) -> requests.Session:
return self._session

@property
def query_builder(self):
if not self._query_builder:
self._query_builder = QueryBuilderWithCache()
def delta_catalog(self) -> DeltaCatalog:
"""The DeltaCatalog used for delta-backed queries, created on first use."""
if self._delta_catalog is None:
self._delta_catalog = DeltaCatalog()
return self._delta_catalog

@property
def query_builder(self) -> QueryBuilderWithCache:
"""Deprecated: use `delta_catalog`."""
warnings.warn(
"`query_builder` is deprecated, use `delta_catalog` instead.",
category=DeprecationWarning,
stacklevel=2,
)
if self._query_builder is None:
self._query_builder = QueryBuilderWithCache(
catalog=self.delta_catalog, _warn=False
)
return self._query_builder

@staticmethod
Expand Down Expand Up @@ -336,6 +370,7 @@ def __init__(
) = MAPI_CLIENT_SETTINGS.LOCAL_DATASET_CACHE,
force_renew: bool = False,
query_builder: QueryBuilderWithCache | None = None,
delta_catalog: DeltaCatalog | None = None,
s3_client: Any | None = None,
timeout: int = 20,
**kwargs,
Expand Down Expand Up @@ -375,9 +410,11 @@ def __init__(
local_dataset_cache: Target directory for downloading full datasets. Defaults
to 'mp_datasets' in the user's home directory
force_renew: Option to overwrite existing local dataset
query_builder : Instance of QueryBuilderWithCache to use in querying delta tables
query_builder : DEPRECATED, use `delta_catalog`. Instance of QueryBuilderWithCache
whose catalog is used for querying delta tables.
NOTE: Must be a QueryBuilderWithCache, a deltalake.QueryBuilder will be ignored.
s3_client: boto3 S3 client object with which to connect to the object stores.ct to the object stores.ct to the object stores.
delta_catalog : Instance of DeltaCatalog to use for querying delta tables.
s3_client: boto3 S3 client object with which to connect to the object stores.
timeout: Time in seconds to wait until a request timeout error is thrown
**kwargs: access to legacy kwargs that may be in the process of being deprecated
"""
Expand All @@ -393,6 +430,7 @@ def __init__(
local_dataset_cache=local_dataset_cache,
force_renew=force_renew,
query_builder=query_builder,
delta_catalog=delta_catalog,
**kwargs,
)

Expand Down Expand Up @@ -594,22 +632,24 @@ def _get_delta_table(
prefix: str,
connector: str = "s3a",
label: str | None = None,
refresh: bool = False,
) -> tuple[str, DeltaTable]:
"""Either create a new DeltaTable, or retrieve a cached one.

If creating a new DeltaTable, will also register in self.query_builder
If creating a new DeltaTable, will also register it in self.delta_catalog

Args:
bucket (str) : name of the bucket in S3
prefix (str) : name of the prefix in S3
connector (str) : s3, s3n, s3a (default), or other
valid Hadoop connector string.
label (str or None) : optional label for the table in the
cached query builder
If `None`, will be gleaned from the URI
label (str or None) : optional label (SQL table name) for the
table in the catalog. If `None`, will be gleaned from the URI
refresh (bool) : if the table is already cached, reload its
snapshot to the latest version first

Returns:
str : the table name in the stored query builder
str : the table name in the catalog
DeltaTable : If one exists at the specified bucket / prefix,
will retrieve the cached instance.
"""
Expand All @@ -621,45 +661,39 @@ def _get_delta_table(
if not uri.endswith("/"):
uri += "/"

try:
stored_label, delta_table = next(
(_label, _table)
for _label, _table in self.query_builder._delta_tables.items()
if _table.table_uri == uri
)
except StopIteration:
stored_label = None

if stored_label is None:
delta_table = DeltaTable(
uri,
storage_options={
"AWS_SKIP_SIGNATURE": "true",
"AWS_REGION": "us-east-1",
"timeout": delta_timeout,
"connect_timeout": delta_timeout,
"pool_idle_timeout": delta_timeout,
"retry_delay": "3",
"max_retries": f"{MAPI_CLIENT_SETTINGS.MAX_RETRIES}",
},
)
self.query_builder.register(qb_label, delta_table)
stored_label, delta_table = self.delta_catalog.get_table(
uri,
qb_label,
storage_options={
"AWS_SKIP_SIGNATURE": "true",
"AWS_REGION": "us-east-1",
"timeout": delta_timeout,
"connect_timeout": delta_timeout,
"pool_idle_timeout": delta_timeout,
"retry_delay": "3",
"max_retries": f"{MAPI_CLIENT_SETTINGS.MAX_RETRIES}",
},
refresh=refresh,
)

elif stored_label != qb_label:
if stored_label != qb_label:
warnings.warn(
f"DeltaTable with URI {uri} already found with different label: "
f"Stored label = {stored_label}; submitted label {qb_label}. "
"Using stored DeltaTable.",
category=MPRestWarning,
stacklevel=2,
)
return stored_label, delta_table

return qb_label, delta_table
return stored_label, delta_table

def _query_delta_single(self, query: str) -> pa.Table:
def _query_delta_single(self, query: str, label: str | None = None) -> pa.Table:
"""Execute a SQL query against a registered Delta table.

If `label` is given and the query fails because a file in the cached
snapshot no longer exists (e.g. the remote table was vacuumed), only
that table is reloaded and the query is retried once.

Wraps the query execution in a try/except to provide a more
actionable error message when the underlying Delta query engine
fails (e.g., due to network timeouts, missing tables, or
Expand All @@ -668,6 +702,8 @@ def _query_delta_single(self, query: str) -> pa.Table:
Args:
query (str): A SQL query string compatible with the
QueryBuilder engine.
label (str or None): The registered table the query reads from,
as returned by `_get_delta_table`. Required for retries.

Returns:
pa.Table: The query result as a PyArrow Table.
Expand All @@ -679,13 +715,20 @@ def _query_delta_single(self, query: str) -> pa.Table:
the underlying cause.
"""
try:
return pa.table(self.query_builder.execute(query).read_all())
return self.delta_catalog.execute(query, label=label)
except Exception as e:
raise MPRestError(
f"Failed to retrieve object due to: {e}. "
f"If this is a timeout error, try increasing the 'timeout' "
f"parameter on MPRester (current value: {self.timeout}s)."
) from e
refreshed = any(
"after refreshing" in note for note in getattr(e, "__notes__", [])
)
hint = (
f"The DeltaTable '{label}' was refreshed and the query retried once."
if refreshed
else (
"If this is a timeout error, try increasing the 'timeout' "
f"parameter on MPRester (current value: {self.timeout}s)."
)
)
raise MPRestError(f"Failed to retrieve object due to: {e}. {hint}") from e

def _query_delta_backed(
self,
Expand Down Expand Up @@ -764,7 +807,8 @@ def _query_delta_backed(
)
}

tbl_lbl, tbl = self._get_delta_table(bucket, prefix, label=label)
# Full downloads are one-off, always start from the latest snapshot
tbl_lbl, tbl = self._get_delta_table(bucket, prefix, label=label, refresh=True)

controlled_batch_str = ",".join(
[f"'{tag}'" for tag in self.access_controlled_batch_ids]
Expand Down Expand Up @@ -810,7 +854,9 @@ def _query_delta_backed(
else None
)

iterator = self.query_builder.execute(f"SELECT * FROM {tbl_lbl} {predicate}")
iterator = self.delta_catalog.execute_stream(
f"SELECT * FROM {tbl_lbl} {predicate}"
)

file_options = ds.ParquetFileFormat().make_write_options(compression="zstd")

Expand Down Expand Up @@ -1749,7 +1795,7 @@ def __getattr__(self, v: str):
db_version=self.db_version,
local_dataset_cache=self.local_dataset_cache,
force_renew=self.force_renew,
query_builder=self._query_builder,
delta_catalog=self.delta_catalog,
)
return self.sub_resters[v]
raise AttributeError(f"{self.__class__} has no attribute {v}")
Expand Down
Loading
Loading