diff --git a/src/datachain/catalog/catalog.py b/src/datachain/catalog/catalog.py index d0eba0962..cd3019ef8 100644 --- a/src/datachain/catalog/catalog.py +++ b/src/datachain/catalog/catalog.py @@ -49,6 +49,7 @@ DatasetInvalidVersionError, DatasetNotFoundError, DatasetVersionNotFoundError, + NamespaceNotFoundError, ProjectNotFoundError, QueryScriptCancelError, QueryScriptRunError, @@ -1107,21 +1108,26 @@ def get_dataset_with_remote_fallback( namespace_name: str, project_name: str, version: Optional[str] = None, + pull_dataset: bool = False, + update: bool = False, ) -> DatasetRecord: - try: - project = self.metastore.get_project(project_name, namespace_name) - ds = self.get_dataset(name, project) - if version and not ds.has_version(version): - raise DatasetVersionNotFoundError( - f"Dataset {name} does not have version {version}" - ) - return ds + if self.metastore.is_local_dataset(namespace_name) or not update: + try: + project = self.metastore.get_project(project_name, namespace_name) + ds = self.get_dataset(name, project) + if not version or ds.has_version(version): + return ds + except (NamespaceNotFoundError, ProjectNotFoundError, DatasetNotFoundError): + pass + + if self.metastore.is_local_dataset(namespace_name): + raise DatasetNotFoundError( + f"Dataset {name}" + + (f" version {version} " if version else " ") + + "not found" + ) - except ( - ProjectNotFoundError, - DatasetNotFoundError, - DatasetVersionNotFoundError, - ): + if pull_dataset: print("Dataset not found in local catalog, trying to get from studio") remote_ds_uri = create_dataset_uri( name, namespace_name, project_name, version @@ -1136,6 +1142,8 @@ def get_dataset_with_remote_fallback( name, self.metastore.get_project(project_name, namespace_name) ) + return self.get_remote_dataset(namespace_name, project_name, name) + def get_dataset_with_version_uuid(self, uuid: str) -> DatasetRecord: """Returns dataset that contains version with specific uuid""" for dataset in self.ls_datasets(): @@ -1152,6 +1160,10 @@ def get_remote_dataset( info_response = studio_client.dataset_info(namespace, project, name) if not info_response.ok: + if info_response.status == 404: + raise DatasetNotFoundError( + f"Dataset {namespace}.{project}.{name} not found" + ) raise DataChainError(info_response.message) dataset_info = info_response.data diff --git a/src/datachain/dataset.py b/src/datachain/dataset.py index a975af17e..085e82903 100644 --- a/src/datachain/dataset.py +++ b/src/datachain/dataset.py @@ -12,6 +12,9 @@ ) from urllib.parse import urlparse +from packaging.specifiers import SpecifierSet +from packaging.version import Version + from datachain import semver from datachain.error import DatasetVersionNotFoundError, InvalidDatasetNameError from datachain.namespace import Namespace @@ -661,13 +664,39 @@ def latest_major_version(self, major: int) -> Optional[str]: return None return max(versions).version - @property - def prev_version(self) -> Optional[str]: - """Returns previous version of a dataset""" - if len(self.versions) == 1: + def latest_compatible_version(self, version_spec: str) -> Optional[str]: + """ + Returns the latest version that matches the given version specifier. + + Supports Python version specifiers like: + - ">=1.0.0,<2.0.0" (compatible release range) + - "~=1.4.2" (compatible release clause) + - "==1.2.*" (prefix matching) + - ">1.0.0" (exclusive ordered comparison) + - ">=1.0.0" (inclusive ordered comparison) + - "!=1.3.0" (version exclusion) + + Args: + version_spec: Version specifier string following PEP 440 + + Returns: + Latest compatible version string, or None if no compatible version found + """ + spec_set = SpecifierSet(version_spec) + + # Convert dataset versions to packaging.Version objects + # and filter compatible ones + compatible_versions = [] + for v in self.versions: + pkg_version = Version(v.version) + if spec_set.contains(pkg_version): + compatible_versions.append(v) + + if not compatible_versions: return None - return sorted(self.versions)[-2].version + # Return the latest compatible version + return max(compatible_versions).version @classmethod def from_dict(cls, d: dict[str, Any]) -> "DatasetRecord": diff --git a/src/datachain/lib/dc/datasets.py b/src/datachain/lib/dc/datasets.py index bd9405755..f7bde7dc2 100644 --- a/src/datachain/lib/dc/datasets.py +++ b/src/datachain/lib/dc/datasets.py @@ -7,9 +7,6 @@ ProjectNotFoundError, ) from datachain.lib.dataset_info import DatasetInfo -from datachain.lib.file import ( - File, -) from datachain.lib.projects import get as get_project from datachain.lib.settings import Settings from datachain.lib.signal_schema import SignalSchema @@ -34,7 +31,6 @@ def read_dataset( version: Optional[Union[str, int]] = None, session: Optional[Session] = None, settings: Optional[dict] = None, - fallback_to_studio: bool = True, delta: Optional[bool] = False, delta_on: Optional[Union[str, Sequence[str]]] = ( "file.path", @@ -44,6 +40,7 @@ def read_dataset( delta_result_on: Optional[Union[str, Sequence[str]]] = None, delta_compare: Optional[Union[str, Sequence[str]]] = None, delta_retry: Optional[Union[bool, str]] = None, + update: bool = False, ) -> "DataChain": """Get data from a saved Dataset. It returns the chain itself. If dataset or version is not found locally, it will try to pull it from Studio. @@ -55,11 +52,12 @@ def read_dataset( set; otherwise, default values will be applied. namespace : optional name of namespace in which dataset to read is created project : optional name of project in which dataset to read is created - version : dataset version + version : dataset version. Supports: + - Exact version strings: "1.2.3" + - Legacy integer versions: 1, 2, 3 (finds latest major version) + - Version specifiers (PEP 440): ">=1.0.0,<2.0.0", "~=1.4.2", "==1.2.*", etc. session : Session to use for the chain. settings : Settings to use for the chain. - fallback_to_studio : Try to pull dataset from Studio if not found locally. - Default is True. delta: If True, only process new or changed files instead of reprocessing everything. This saves time by skipping files that were already processed in previous versions. The optimization is working when a new version of the @@ -79,6 +77,10 @@ def read_dataset( (error mode) - True: Reprocess records missing from the result dataset (missing mode) - None: No retry processing (default) + update: If True always checks for newer versions available on Studio, even if + some version of the dataset exists locally already. If False (default), it + will only fetch the dataset from Studio if it is not found locally. + Example: ```py @@ -92,11 +94,22 @@ def read_dataset( ``` ```py - chain = dc.read_dataset("my_cats", fallback_to_studio=False) + chain = dc.read_dataset("my_cats", version="1.0.0") ``` ```py - chain = dc.read_dataset("my_cats", version="1.0.0") + # Using version specifiers (PEP 440) + chain = dc.read_dataset("my_cats", version=">=1.0.0,<2.0.0") + ``` + + ```py + # Legacy integer version support (finds latest in major version) + chain = dc.read_dataset("my_cats", version=1) # Latest 1.x.x version + ``` + + ```py + # Always check for newer versions matching a version specifier from Studio + chain = dc.read_dataset("my_cats", version=">=1.0.0", update=True) ``` ```py @@ -113,7 +126,6 @@ def read_dataset( version="1.0.0", session=session, settings=settings, - fallback_to_studio=True, ) ``` """ @@ -121,6 +133,8 @@ def read_dataset( from .datachain import DataChain + telemetry.send_event_once("class", "datachain_init", name=name, version=version) + session = Session.get(session) catalog = session.catalog @@ -131,31 +145,37 @@ def read_dataset( ) if version is not None: + dataset = session.catalog.get_dataset_with_remote_fallback( + name, namespace_name, project_name, update=update + ) + + # Convert legacy integer versions to version specifiers + # For backward compatibility we still allow users to put version as integer + # in which case we convert it to a version specifier that finds the latest + # version where major part is equal to that input version. + # For example if user sets version=2, we convert it to ">=2.0.0,<3.0.0" + # which will find something like 2.4.3 (assuming 2.4.3 is the biggest among + # all 2.* dataset versions) + if isinstance(version, int): + version_spec = f">={version}.0.0,<{version + 1}.0.0" + else: + version_spec = str(version) + + from packaging.specifiers import InvalidSpecifier, SpecifierSet + try: - # for backward compatibility we still allow users to put version as integer - # in which case we are trying to find latest version where major part is - # equal to that input version. For example if user sets version=2, we could - # continue with something like 2.4.3 (assuming 2.4.3 is the biggest among - # all 2.* dataset versions). If dataset doesn't have any versions where - # major part is equal to that input, exception is thrown. - major = int(version) - try: - ds_project = get_project(project_name, namespace_name, session=session) - except ProjectNotFoundError: - raise DatasetNotFoundError( - f"Dataset {name} not found in namespace {namespace_name} and", - f" project {project_name}", - ) from None - - dataset = session.catalog.get_dataset(name, ds_project) - latest_major = dataset.latest_major_version(major) - if not latest_major: + # Try to parse as version specifier + SpecifierSet(version_spec) + # If it's a valid specifier set, find the latest compatible version + latest_compatible = dataset.latest_compatible_version(version_spec) + if not latest_compatible: raise DatasetVersionNotFoundError( - f"Dataset {name} does not have version {version}" + f"No dataset {name} version matching specifier {version_spec}" ) - version = latest_major - except ValueError: - # version is in new semver string format, continuing as normal + version = latest_compatible + except InvalidSpecifier: + # If not a valid specifier, treat as exact version string + # This handles cases like "1.2.3" which are exact versions, not specifiers pass if settings: @@ -169,11 +189,8 @@ def read_dataset( namespace_name=namespace_name, version=version, # type: ignore[arg-type] session=session, - indexing_column_types=File._datachain_column_types, - fallback_to_studio=fallback_to_studio, ) - telemetry.send_event_once("class", "datachain_init", name=name, version=version) signals_schema = SignalSchema({"sys": Sys}) if query.feature_schema: signals_schema |= SignalSchema.deserialize(query.feature_schema) diff --git a/src/datachain/lib/dc/listings.py b/src/datachain/lib/dc/listings.py index 594caeff7..71d25eea5 100644 --- a/src/datachain/lib/dc/listings.py +++ b/src/datachain/lib/dc/listings.py @@ -127,12 +127,8 @@ def read_listing_dataset( if version is None: version = dataset.latest_version - query = DatasetQuery( - name=name, - session=session, - indexing_column_types=File._datachain_column_types, - fallback_to_studio=False, - ) + query = DatasetQuery(name=name, session=session) + if settings: cfg = {**settings} if "prefetch" not in cfg: diff --git a/src/datachain/lib/projects.py b/src/datachain/lib/projects.py index 942ddc86b..f14385fc2 100644 --- a/src/datachain/lib/projects.py +++ b/src/datachain/lib/projects.py @@ -54,7 +54,7 @@ def get(name: str, namespace: str, session: Optional[Session]) -> Project: ```py import datachain as dc from datachain.lib.projects import get as get_project - project = get_project("my-project", "local") + project = get_project("my-project", "local") ``` """ return Session.get(session).catalog.metastore.get_project(name, namespace) diff --git a/src/datachain/query/dataset.py b/src/datachain/query/dataset.py index 1d3e2d589..2508675e0 100644 --- a/src/datachain/query/dataset.py +++ b/src/datachain/query/dataset.py @@ -1099,13 +1099,9 @@ def __init__( namespace_name: Optional[str] = None, catalog: Optional["Catalog"] = None, session: Optional[Session] = None, - indexing_column_types: Optional[dict[str, Any]] = None, in_memory: bool = False, - fallback_to_studio: bool = True, update: bool = False, ) -> None: - from datachain.remote.studio import is_token_set - self.session = Session.get(session, catalog=catalog, in_memory=in_memory) self.catalog = catalog or self.session.catalog self.steps: list[Step] = [] @@ -1137,18 +1133,16 @@ def __init__( # not setting query step yet as listing dataset might not exist at # this point self.list_ds_name = name - elif fallback_to_studio and is_token_set(): + else: self._set_starting_step( self.catalog.get_dataset_with_remote_fallback( name, namespace_name=namespace_name, project_name=project_name, version=version, + pull_dataset=True, ) ) - else: - project = self.catalog.metastore.get_project(project_name, namespace_name) - self._set_starting_step(self.catalog.get_dataset(name, project=project)) def _set_starting_step(self, ds: "DatasetRecord") -> None: if not self.version: diff --git a/src/datachain/remote/studio.py b/src/datachain/remote/studio.py index c951653da..1824ab988 100644 --- a/src/datachain/remote/studio.py +++ b/src/datachain/remote/studio.py @@ -78,10 +78,11 @@ def _parse_dates(obj: dict, date_fields: list[str]): class Response(Generic[T]): - def __init__(self, data: T, ok: bool, message: str) -> None: + def __init__(self, data: T, ok: bool, message: str, status: int) -> None: self.data = data self.ok = ok self.message = message + self.status = status def __repr__(self): return ( @@ -186,7 +187,7 @@ def _send_request_msgpack( message = "Indexing in progress" else: message = content.get("message", "") - return Response(response_data, ok, message) + return Response(response_data, ok, message, response.status_code) @retry_with_backoff(retries=3, errors=(HTTPError, Timeout)) def _send_request( @@ -236,7 +237,7 @@ def _send_request( else: message = "" - return Response(data, ok, message) + return Response(data, ok, message, response.status_code) @staticmethod def _unpacker_hook(code, data): diff --git a/tests/conftest.py b/tests/conftest.py index 9b32c9fa5..59b4aeb61 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -904,3 +904,172 @@ def run_datachain_worker(datachain_job_id): worker.wait(timeout=30) # seconds except subprocess.TimeoutExpired: os.kill(worker.pid, signal.SIGKILL) + + +# Common constants for remote dataset testing +REMOTE_DATASET_UUID = "20f5a2f1-fc9a-4e36-8b91-5a530f289451" +REMOTE_DATASET_UUID_V2 = "30f5a2f1-fc9a-4e36-8b91-5a530f289452" +REMOTE_NAMESPACE_UUID = "efbc3e84-3eb6-4be1-bec1-704b939e1fe4" +REMOTE_PROJECT_UUID = "0ed3a6c6-f2f7-45aa-869b-39219c86a9d4" + +REMOTE_NAMESPACE_NAME = "dev" +REMOTE_PROJECT_NAME = "animals" + + +@pytest.fixture +def remote_dataset_schema(): + """Common schema for remote datasets.""" + return { + "id": {"type": "UInt64"}, + "sys__rand": {"type": "UInt64"}, + "file__path": {"type": "String"}, + "file__etag": {"type": "String"}, + "file__version": {"type": "String"}, + "file__is_latest": {"type": "Boolean"}, + "file__last_modified": {"type": "DateTime"}, + "file__size": {"type": "Int64"}, + "file__location": {"type": "String"}, + "file__source": {"type": "String"}, + "version": {"type": "String"}, + } + + +@pytest.fixture +def remote_file_feature_schema(): + """Common File feature schema for remote datasets.""" + return { + "file": "File@v1", + "version": "str", + "_custom_types": { + "File@v1": { + "schema_version": 2, + "name": "File@v1", + "fields": { + "source": "str", + "path": "str", + "size": "int", + "version": "str", + "etag": "str", + "is_latest": "bool", + "last_modified": "datetime", + "location": "Union[dict, list[dict], NoneType]", + }, + "bases": [ + ["File", "datachain.lib.file", "File@v1"], + ["DataModel", "datachain.lib.data_model", "DataModel@v1"], + ["BaseModel", "pydantic.main", None], + ["object", "builtins", None], + ], + "hidden_fields": [ + "source", + "version", + "etag", + "is_latest", + "last_modified", + "location", + ], + } + }, + } + + +@pytest.fixture +def remote_namespace(): + """Remote namespace fixture for Studio API mocking.""" + return { + "id": 1, + "uuid": REMOTE_NAMESPACE_UUID, + "name": REMOTE_NAMESPACE_NAME, + "descr": "Dev namespace", + "created_at": "2024-02-23T10:42:31.842944+00:00", + } + + +@pytest.fixture +def remote_project(remote_namespace): + """Remote project fixture for Studio API mocking.""" + return { + "id": 1, + "uuid": REMOTE_PROJECT_UUID, + "name": REMOTE_PROJECT_NAME, + "descr": "Animals project", + "created_at": "2024-02-23T10:42:31.842944+00:00", + "namespace": remote_namespace, + } + + +@pytest.fixture +def compressed_parquet_data(): + """ + Factory fixture that creates lz4-compressed parquet for datasets. + Returns a function that can be called with different data. + """ + import io + from datetime import datetime + + import lz4.frame + import pandas as pd + + def create_compressed_parquet(data, src_uri=None): + def _adapt_row(row): + """ + Adjusting row values to match remote response + """ + adapted = {} + for k, v in row.items(): + if isinstance(v, datetime): + adapted[k] = v.timestamp() + elif v is None: + adapted[k] = "" + else: + adapted[k] = v + + adapted["sys__id"] = 1 + adapted["sys__rand"] = 1 + adapted["file__location"] = "" + adapted["file__source"] = src_uri or "" + adapted["file__version"] = "" + return adapted + + adapted_data = [_adapt_row(row) for row in data] + df = pd.DataFrame.from_records(adapted_data) + buffer = io.BytesIO() + df.to_parquet(buffer, engine="auto") + + return lz4.frame.compress(buffer.getvalue()) + + return create_compressed_parquet + + +@pytest.fixture +def dog_entries(): + """Factory function to create version-specific dog entries.""" + from tests.data import ENTRIES + + def _create_dog_entries(version="1.0.0"): + return [ + { + "file__path": e.path, + "file__etag": e.etag, + "file__version": e.version, + "file__is_latest": e.is_latest, + "file__last_modified": e.last_modified, + "file__size": e.size, + "version": version, + } + for e in ENTRIES + if e.name.startswith("dog") + ] + + return _create_dog_entries + + +@pytest.fixture +def mock_parquet_data(compressed_parquet_data, dog_entries, version="1.0.0"): + return compressed_parquet_data(dog_entries(version)) + + +@pytest.fixture +def mock_parquet_data_cloud(compressed_parquet_data, dog_entries, cloud_test_catalog): + src_uri = cloud_test_catalog.src_uri + return compressed_parquet_data(dog_entries("1.0.0"), src_uri) diff --git a/tests/func/test_dataset_query.py b/tests/func/test_dataset_query.py index 62a83fb9f..59a991d4b 100644 --- a/tests/func/test_dataset_query.py +++ b/tests/func/test_dataset_query.py @@ -7,9 +7,7 @@ import sqlalchemy from datachain.dataset import DatasetDependencyType, DatasetStatus -from datachain.error import ( - DatasetVersionNotFoundError, -) +from datachain.error import DatasetNotFoundError from datachain.lib.listing import parse_listing_uri from datachain.query import C, DatasetQuery, Object, Stream from datachain.sql.functions import path as pathfunc @@ -70,7 +68,7 @@ def test_save_multiple_versions(cloud_test_catalog, animal_dataset): assert DatasetQuery(name=ds_name, version="1.0.1", catalog=catalog).count() == 3 assert DatasetQuery(name=ds_name, version="1.0.2", catalog=catalog).count() == 3 - with pytest.raises(DatasetVersionNotFoundError): + with pytest.raises(DatasetNotFoundError): DatasetQuery(name=ds_name, version="4.0.0", catalog=catalog).count() diff --git a/tests/func/test_delta.py b/tests/func/test_delta.py index 921af57a2..f88058318 100644 --- a/tests/func/test_delta.py +++ b/tests/func/test_delta.py @@ -6,7 +6,7 @@ import datachain as dc from datachain import func -from datachain.error import DatasetVersionNotFoundError +from datachain.error import DatasetNotFoundError from datachain.lib.dc import C from datachain.lib.file import File, ImageFile @@ -278,10 +278,10 @@ def get_index(file: File) -> int: "images/img9.jpg", ] - with pytest.raises(DatasetVersionNotFoundError) as exc_info: + with pytest.raises(DatasetNotFoundError) as exc_info: dc.read_dataset(ds_name, version="1.0.1") - assert str(exc_info.value) == f"Dataset {ds_name} does not have version 1.0.1" + assert str(exc_info.value) == f"Dataset {ds_name} version 1.0.1 not found" @pytest.fixture diff --git a/tests/func/test_ls.py b/tests/func/test_ls.py index 3ded496c3..c31460219 100644 --- a/tests/func/test_ls.py +++ b/tests/func/test_ls.py @@ -179,9 +179,10 @@ def test_ls_partial_indexing(cloud_test_catalog, cloud_type, capsys): class MockResponse: - def __init__(self, content, ok=True): + def __init__(self, content, status_code, ok=True): self.content = content self.ok = ok + self.status_code = status_code def mock_post(method, url, data=None, json=None, **kwargs): @@ -196,7 +197,8 @@ def mock_post(method, url, data=None, json=None, **kwargs): for d in REMOTE_DATA[path] ] return MockResponse( - content=msgpack.packb({"data": data}, default=_pack_extended_types) + content=msgpack.packb({"data": data}, default=_pack_extended_types), + status_code=200, ) diff --git a/tests/func/test_pull.py b/tests/func/test_pull.py index 29bf6a2ae..d0b0e9720 100644 --- a/tests/func/test_pull.py +++ b/tests/func/test_pull.py @@ -1,9 +1,5 @@ -import io import json -from datetime import datetime -import lz4.frame -import pandas as pd import pytest import datachain as dc @@ -12,126 +8,27 @@ from datachain.error import DataChainError, DatasetNotFoundError from datachain.query.session import Session from datachain.utils import STUDIO_URL, JSONSerialize -from tests.data import ENTRIES +from tests.conftest import ( + REMOTE_DATASET_UUID, + REMOTE_NAMESPACE_NAME, + REMOTE_NAMESPACE_UUID, + REMOTE_PROJECT_NAME, + REMOTE_PROJECT_UUID, +) from tests.utils import assert_row_names, skip_if_not_sqlite, tree_from_path -DATASET_UUID = "20f5a2f1-fc9a-4e36-8b91-5a530f289451" -NAMESPACE_UUID = "efbc3e84-3eb6-4be1-bec1-704b939e1fe4" -PROJECT_UUID = "0ed3a6c6-f2f7-45aa-869b-39219c86a9d4" - -NAMESPACE_NAME = "dev" -PROJECT_NAME = "animals" - - -@pytest.fixture -def dog_entries(): - # TODO remove when we replace ENTRIES with FILES - return [ - { - "file__path": e.path, - "file__etag": e.etag, - "file__version": e.version, - "file__is_latest": e.is_latest, - "file__last_modified": e.last_modified, - "file__size": e.size, - } - for e in ENTRIES - if e.name.startswith("dog") - ] - - -@pytest.fixture -def dog_entries_parquet_lz4(dog_entries, cloud_test_catalog) -> bytes: - """ - Returns dogs entries in lz4 compressed parquet format - """ - src_uri = cloud_test_catalog.src_uri - - def _adapt_row(row): - """ - Adjusting row values to match remote response - """ - adapted = {} - for k, v in row.items(): - if isinstance(v, datetime): - adapted[k] = v.timestamp() - elif v is None: - adapted[k] = "" - else: - adapted[k] = v - - adapted["sys__id"] = 1 - adapted["sys__rand"] = 1 - adapted["file__location"] = "" - adapted["file__source"] = src_uri - adapted["file__version"] = "" - return adapted - - dog_entries = [_adapt_row(e) for e in dog_entries] - df = pd.DataFrame.from_records(dog_entries) - buffer = io.BytesIO() - df.to_parquet(buffer, engine="auto") - - return lz4.frame.compress(buffer.getvalue()) - @pytest.fixture -def schema(): - return { - "id": {"type": "UInt64"}, - "sys__rand": {"type": "UInt64"}, - "file__path": {"type": "String"}, - "file__etag": {"type": "String"}, - "file__version": {"type": "String"}, - "file__is_latest": {"type": "Boolean"}, - "file__last_modified": {"type": "DateTime"}, - "file__size": {"type": "Int64"}, - "file__location": {"type": "String"}, - "file__source": {"type": "String"}, - } - - -@pytest.fixture -def remote_dataset_version(schema, dataset_rows): +def remote_dataset_version( + remote_dataset_schema, dataset_rows, remote_file_feature_schema +): return { "id": 1, - "uuid": DATASET_UUID, + "uuid": REMOTE_DATASET_UUID, "dataset_id": 1, "version": "1.0.0", "status": 4, - "feature_schema": { - "file": "File@v1", - "_custom_types": { - "File@v1": { - "schema_version": 2, - "name": "File@v1", - "fields": { - "source": "str", - "path": "str", - "size": "int", - "version": "str", - "etag": "str", - "is_latest": "bool", - "last_modified": "datetime", - "location": "Union[dict, list[dict], NoneType]", - }, - "bases": [ - ["File", "datachain.lib.file", "File@v1"], - ["DataModel", "datachain.lib.data_model", "DataModel@v1"], - ["BaseModel", "pydantic.main", None], - ["object", "builtins", None], - ], - "hidden_fields": [ - "source", - "version", - "etag", - "is_latest", - "last_modified", - "location", - ], - } - }, - }, + "feature_schema": remote_file_feature_schema, "created_at": "2024-02-23T10:42:31.842944+00:00", "finished_at": "2024-02-23T10:43:31.842944+00:00", "error_message": "", @@ -140,7 +37,7 @@ def remote_dataset_version(schema, dataset_rows): "size": 1073741824, "preview": json.loads(json.dumps(dataset_rows, cls=JSONSerialize)), "script_output": "", - "schema": schema, + "schema": remote_dataset_schema, "sources": "", "query_script": ( 'from datachain.query.dataset import DatasetQuery\nDatasetQuery(path="s3://ldb-public")', @@ -150,71 +47,21 @@ def remote_dataset_version(schema, dataset_rows): @pytest.fixture -def remote_namespace(): - return { - "id": 1, - "uuid": NAMESPACE_UUID, - "name": NAMESPACE_NAME, - "descr": "Dev namespace", - "created_at": "2024-02-23T10:42:31.842944+00:00", - } - - -@pytest.fixture -def remote_project(remote_namespace): - return { - "id": 1, - "uuid": PROJECT_UUID, - "name": PROJECT_NAME, - "descr": "Animals project", - "created_at": "2024-02-23T10:42:31.842944+00:00", - "namespace": remote_namespace, - } - - -@pytest.fixture -def remote_dataset(remote_project, remote_dataset_version, schema): +def remote_dataset( + remote_project, + remote_dataset_version, + remote_dataset_schema, + remote_file_feature_schema, +): return { "id": 1, "name": "dogs", "project": remote_project, "description": "", "attrs": [], - "schema": schema, + "schema": remote_dataset_schema, "status": 4, - "feature_schema": { - "file": "File@v1", - "_custom_types": { - "File@v1": { - "schema_version": 2, - "name": "File@v1", - "fields": { - "source": "str", - "path": "str", - "size": "int", - "version": "str", - "etag": "str", - "is_latest": "bool", - "last_modified": "datetime", - "location": "Union[dict, list[dict], NoneType]", - }, - "bases": [ - ["File", "datachain.lib.file", "File@v1"], - ["DataModel", "datachain.lib.data_model", "DataModel@v1"], - ["BaseModel", "pydantic.main", None], - ["object", "builtins", None], - ], - "hidden_fields": [ - "source", - "version", - "etag", - "is_latest", - "last_modified", - "location", - ], - } - }, - }, + "feature_schema": remote_file_feature_schema, "created_at": "2024-02-23T10:42:31.842944+00:00", "finished_at": "2024-02-23T10:43:31.842944+00:00", "error_message": "", @@ -259,17 +106,17 @@ def dataset_export_status(requests_mock): @pytest.fixture def dataset_export_data_chunk( - requests_mock, remote_dataset_chunk_url, dog_entries_parquet_lz4 + requests_mock, remote_dataset_chunk_url, mock_parquet_data_cloud ): - requests_mock.get(remote_dataset_chunk_url, content=dog_entries_parquet_lz4) + requests_mock.get(remote_dataset_chunk_url, content=mock_parquet_data_cloud) @pytest.mark.parametrize("cloud_type, version_aware", [("s3", False)], indirect=True) @pytest.mark.parametrize( "dataset_uri", [ - f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0", - f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs", + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0", + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", ], ) @pytest.mark.parametrize("local_ds_name", [None, "other"]) @@ -320,10 +167,10 @@ def test_pull_dataset_success( cp=False, ) - project = catalog.metastore.get_project(PROJECT_NAME, NAMESPACE_NAME) + project = catalog.metastore.get_project(REMOTE_PROJECT_NAME, REMOTE_NAMESPACE_NAME) dataset = catalog.get_dataset(local_ds_name or "dogs", project=project) - assert dataset.project.namespace.uuid == NAMESPACE_UUID - assert dataset.project.uuid == PROJECT_UUID + assert dataset.project.namespace.uuid == REMOTE_NAMESPACE_UUID + assert dataset.project.uuid == REMOTE_PROJECT_UUID assert [v.version for v in dataset.versions] == [local_ds_version or "1.0.0"] assert dataset.status == DatasetStatus.COMPLETE @@ -337,7 +184,7 @@ def test_pull_dataset_success( assert dataset_version.schema assert dataset_version.num_objects == 4 assert dataset_version.size == 15 - assert dataset_version.uuid == DATASET_UUID + assert dataset_version.uuid == REMOTE_DATASET_UUID assert_row_names( catalog, @@ -391,9 +238,8 @@ def test_datachain_read_dataset_pull( with Session("testSession", catalog=catalog): ds = dc.read_dataset( - name=f"{NAMESPACE_NAME}.{PROJECT_NAME}.dogs", + name=f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", version="1.0.0", - fallback_to_studio=True, ) assert ds.dataset.name == "dogs" @@ -401,7 +247,7 @@ def test_datachain_read_dataset_pull( assert ds.dataset.status == DatasetStatus.COMPLETE # Check that dataset is available locally after pulling - project = catalog.metastore.get_project(PROJECT_NAME, NAMESPACE_NAME) + project = catalog.metastore.get_project(REMOTE_PROJECT_NAME, REMOTE_NAMESPACE_NAME) dataset = catalog.get_dataset("dogs", project) assert dataset.name == "dogs" @@ -432,7 +278,9 @@ def test_pull_dataset_wrong_version( catalog = cloud_test_catalog.catalog with pytest.raises(DataChainError) as exc_info: - catalog.pull_dataset(f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v5") + catalog.pull_dataset( + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v5" + ) assert str(exc_info.value) == "Dataset dogs doesn't have version 5 on server" @@ -450,9 +298,13 @@ def test_pull_dataset_not_found_in_remote( ) catalog = cloud_test_catalog.catalog - with pytest.raises(DataChainError) as exc_info: - catalog.pull_dataset(f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0") - assert str(exc_info.value) == "Dataset not found" + full_name = f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs" + + with pytest.raises(DatasetNotFoundError) as exc_info: + catalog.pull_dataset( + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0" + ) + assert str(exc_info.value) == f"Dataset {full_name} not found" @pytest.mark.parametrize("cloud_type, version_aware", [("s3", False)], indirect=True) @@ -474,7 +326,9 @@ def test_pull_dataset_exporting_dataset_failed_in_remote( catalog = cloud_test_catalog.catalog with pytest.raises(DataChainError) as exc_info: - catalog.pull_dataset(f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0") + catalog.pull_dataset( + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0" + ) assert str(exc_info.value) == f"Dataset export {export_status} in Studio" @@ -493,7 +347,9 @@ def test_pull_dataset_empty_parquet( catalog = cloud_test_catalog.catalog with pytest.raises(RuntimeError): - catalog.pull_dataset(f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0") + catalog.pull_dataset( + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0" + ) @pytest.mark.parametrize("cloud_type, version_aware", [("s3", False)], indirect=True) @@ -509,14 +365,17 @@ def test_pull_dataset_already_exists_locally( catalog = cloud_test_catalog.catalog catalog.pull_dataset( - f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0", local_ds_name="other" + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0", + local_ds_name="other", + ) + catalog.pull_dataset( + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0" ) - catalog.pull_dataset(f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0") - project = catalog.metastore.get_project(PROJECT_NAME, NAMESPACE_NAME) + project = catalog.metastore.get_project(REMOTE_PROJECT_NAME, REMOTE_NAMESPACE_NAME) other = catalog.get_dataset("other", project) other_version = other.get_version("1.0.0") - assert other_version.uuid == DATASET_UUID + assert other_version.uuid == REMOTE_DATASET_UUID assert other_version.num_objects == 4 assert other_version.size == 15 @@ -540,25 +399,28 @@ def test_pull_dataset_local_name_already_exists( catalog = cloud_test_catalog.catalog src_uri = cloud_test_catalog.src_uri - project = catalog.metastore.create_project(NAMESPACE_NAME, PROJECT_NAME) + project = catalog.metastore.create_project( + REMOTE_NAMESPACE_NAME, REMOTE_PROJECT_NAME + ) catalog.create_dataset_from_sources( local_ds_name or "dogs", [f"{src_uri}/dogs/*"], recursive=True, project=project ) with pytest.raises(DataChainError) as exc_info: catalog.pull_dataset( - f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0", + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0", local_ds_name=local_ds_name, ) assert str(exc_info.value) == ( - f"Local dataset ds://{NAMESPACE_NAME}.{PROJECT_NAME}.{local_ds_name or 'dogs'}" + f"Local dataset ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}." + f"{local_ds_name or 'dogs'}" "@v1.0.0 already exists with" " different uuid, please choose different local dataset name or version" ) # able to save it as version 2 of local dataset name catalog.pull_dataset( - f"ds://{NAMESPACE_NAME}.{PROJECT_NAME}.dogs@v1.0.0", + f"ds://{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs@v1.0.0", local_ds_name=local_ds_name, local_ds_version="2.0.0", ) diff --git a/tests/func/test_read_dataset_remote.py b/tests/func/test_read_dataset_remote.py new file mode 100644 index 000000000..8e9aaed76 --- /dev/null +++ b/tests/func/test_read_dataset_remote.py @@ -0,0 +1,555 @@ +""" +Functional tests for read_dataset when accessing remote (Studio) datasets. +""" + +import json +from urllib.parse import parse_qs, urlparse + +import pytest + +import datachain as dc +from datachain.error import ( + DataChainError, + DatasetNotFoundError, + DatasetVersionNotFoundError, +) +from datachain.utils import STUDIO_URL, JSONSerialize +from tests.conftest import ( + REMOTE_DATASET_UUID, + REMOTE_DATASET_UUID_V2, + REMOTE_NAMESPACE_NAME, + REMOTE_PROJECT_NAME, +) +from tests.utils import skip_if_not_sqlite + + +@pytest.fixture +def remote_dataset_version_v1( + remote_dataset_schema, dataset_rows, remote_file_feature_schema +): + return { + "id": 1, + "uuid": REMOTE_DATASET_UUID, + "dataset_id": 1, + "version": "1.0.0", + "status": 4, + "feature_schema": remote_file_feature_schema, + "created_at": "2024-02-23T10:42:31.842944+00:00", + "finished_at": "2024-02-23T10:43:31.842944+00:00", + "error_message": "", + "error_stack": "", + "num_objects": 1, + "size": 1024, + "preview": json.loads(json.dumps(dataset_rows, cls=JSONSerialize)), + "script_output": "", + "schema": remote_dataset_schema, + "sources": "", + "query_script": ( + "from datachain.query.dataset import DatasetQuery\n" + 'DatasetQuery(path="s3://test-bucket")', + ), + "created_by_id": 1, + } + + +@pytest.fixture +def remote_dataset_version_v2( + remote_dataset_schema, dataset_rows, remote_file_feature_schema +): + return { + "id": 2, + "uuid": REMOTE_DATASET_UUID_V2, + "dataset_id": 1, + "version": "2.0.0", + "status": 4, + "feature_schema": remote_file_feature_schema, + "created_at": "2024-02-24T10:42:31.842944+00:00", + "finished_at": "2024-02-24T10:43:31.842944+00:00", + "error_message": "", + "error_stack": "", + "num_objects": 1, + "size": 2048, + "preview": json.loads(json.dumps(dataset_rows, cls=JSONSerialize)), + "script_output": "", + "schema": remote_dataset_schema, + "sources": "", + "query_script": ( + "from datachain.query.dataset import DatasetQuery\n" + 'DatasetQuery(path="s3://test-bucket")', + ), + "created_by_id": 1, + } + + +@pytest.fixture +def remote_dataset_single_version( + remote_project, + remote_dataset_version_v1, + remote_dataset_schema, + remote_file_feature_schema, +): + return { + "id": 1, + "name": "dogs", + "project": remote_project, + "description": "", + "attrs": [], + "schema": remote_dataset_schema, + "status": 4, + "feature_schema": remote_file_feature_schema, + "created_at": "2024-02-23T10:42:31.842944+00:00", + "finished_at": "2024-02-23T10:43:31.842944+00:00", + "error_message": "", + "error_stack": "", + "script_output": "", + "job_id": "f74ec414-58b7-437d-81c5-d41e5365abba", + "sources": "", + "query_script": "", + "team_id": 1, + "warehouse_id": None, + "created_by_id": 1, + "versions": [remote_dataset_version_v1], + } + + +@pytest.fixture +def remote_dataset_multi_version( + remote_project, + remote_dataset_version_v1, + remote_dataset_version_v2, + remote_dataset_schema, + remote_file_feature_schema, +): + return { + "id": 1, + "name": "dogs", + "project": remote_project, + "description": "", + "attrs": [], + "schema": remote_dataset_schema, + "status": 4, + "feature_schema": remote_file_feature_schema, + "created_at": "2024-02-23T10:42:31.842944+00:00", + "finished_at": "2024-02-23T10:43:31.842944+00:00", + "error_message": "", + "error_stack": "", + "script_output": "", + "job_id": "f74ec414-58b7-437d-81c5-d41e5365abba", + "sources": "", + "query_script": "", + "team_id": 1, + "warehouse_id": None, + "created_by_id": 1, + "versions": [remote_dataset_version_v1, remote_dataset_version_v2], + } + + +@pytest.fixture +def mock_dataset_info_endpoint(requests_mock): + """Mock the dataset info endpoint to return dataset information.""" + + def _mock_info(dataset_data): + return requests_mock.get( + f"{STUDIO_URL}/api/datachain/datasets/info", json=dataset_data + ) + + return _mock_info + + +@pytest.fixture +def mock_dataset_info_not_found(requests_mock): + """Mock the dataset info endpoint to return 404 not found.""" + return requests_mock.get( + f"{STUDIO_URL}/api/datachain/datasets/info", + status_code=404, + json={"message": "Dataset not found"}, + ) + + +def _get_version_from_request(request, default="1.0.0"): + parsed_url = urlparse(request.url) + query_params = parse_qs(parsed_url.query) + return query_params.get("version", [default])[0] + + +@pytest.fixture +def mock_export_endpoint_with_urls(requests_mock): + """Mock the export endpoint to return download URLs based on version.""" + + def _mock_export_response(request, context): + version_param = _get_version_from_request(request) + version_file = version_param.replace(".", "_") + return [ + f"https://studio-blobvault.s3.amazonaws.com/" + f"datachain_ds_export_{version_file}.parquet.lz4" + ] + + return requests_mock.get( + f"{STUDIO_URL}/api/datachain/datasets/export", json=_mock_export_response + ) + + +@pytest.fixture +def mock_export_endpoint_with_id(requests_mock): + """Mock the export endpoint to return an export ID based on version.""" + + def _mock_export_id_response(request, context): + version_param = _get_version_from_request(request) + return {"export_id": f"test-export-id-{version_param}"} + + return requests_mock.get( + f"{STUDIO_URL}/api/datachain/datasets/export", json=_mock_export_id_response + ) + + +@pytest.fixture +def mock_export_status_completed(requests_mock): + """Mock the export status endpoint to return completed status based on version.""" + + def _mock_status_response(request, context): + version_param = _get_version_from_request(request) + return {"status": "completed", "version": version_param} + + return requests_mock.get( + f"{STUDIO_URL}/api/datachain/datasets/export-status", json=_mock_status_response + ) + + +@pytest.fixture +def mock_export_status_failed(requests_mock): + """Mock the export status endpoint to return failed status based on version.""" + + def _mock_failed_response(request, context): + version_param = _get_version_from_request(request) + return { + "status": "failed", + "version": version_param, + "error": f"Export failed for version {version_param}", + } + + return requests_mock.get( + f"{STUDIO_URL}/api/datachain/datasets/export-status", json=_mock_failed_response + ) + + +@pytest.fixture +def mock_s3_parquet_download(requests_mock, compressed_parquet_data, dog_entries): + """Mock S3 parquet file download for all versions.""" + + def _mock_download(): + # Generate different data for each version + for version in ["1.0.0", "2.0.0"]: + parquet_data = compressed_parquet_data(dog_entries(version)) + requests_mock.get( + f"https://studio-blobvault.s3.amazonaws.com/" + f"datachain_ds_export_{version.replace('.', '_')}.parquet.lz4", + content=parquet_data, + ) + + return _mock_download + + +@pytest.fixture +def mock_dataset_rows_fetcher_status_check(mocker): + """Mock DatasetRowsFetcher.should_check_for_status to return True.""" + return mocker.patch( + "datachain.catalog.catalog.DatasetRowsFetcher.should_check_for_status", + return_value=True, + ) + + +@skip_if_not_sqlite +def test_read_dataset_remote_basic( + studio_token, + test_session, + remote_dataset_single_version, + mock_dataset_info_endpoint, + mock_export_endpoint_with_urls, + mock_export_status_completed, + mock_s3_parquet_download, + mock_dataset_rows_fetcher_status_check, +): + """Test basic read_dataset functionality with remote dataset.""" + mock_dataset_info_endpoint(remote_dataset_single_version) + mock_s3_parquet_download() + + # Ensure dataset is not available locally at first + with pytest.raises(DatasetNotFoundError): + dc.read_dataset("dogs", session=test_session) + + assert ( + dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version="1.0.0", + session=test_session, + ).to_values("version")[0] + == "1.0.0" + ) + + +@skip_if_not_sqlite +def test_read_dataset_remote_already_exists( + studio_token, + test_session, + requests_mock, + remote_dataset_single_version, + mock_dataset_info_endpoint, + mock_export_endpoint_with_urls, + mock_export_status_completed, + mock_s3_parquet_download, + mock_dataset_rows_fetcher_status_check, +): + """Test read_dataset when dataset already exists locally from previous read.""" + + # Mock the Studio API responses + mock_dataset_info_endpoint(remote_dataset_single_version) + mock_s3_parquet_download() + + # First read - downloads from remote + ds1 = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version="1.0.0", + session=test_session, + ) + + assert ds1.to_values("version")[0] == "1.0.0" + assert ds1.dataset.name == "dogs" + assert dc.datasets().to_values("version") == ["1.0.0"] + + # Second read - should use local dataset without calling remote + # Clear the mock to ensure no new remote calls are made + requests_mock.reset_mock() + + ds2 = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version="1.0.0", + session=test_session, + ) + + assert ds2.to_values("version")[0] == "1.0.0" + assert ds2.dataset.name == "dogs" + assert dc.datasets().to_values("version") == ["1.0.0"] + assert ds2.dataset.versions[0].uuid == REMOTE_DATASET_UUID + + # Verify no remote calls were made for the second read + assert not requests_mock.called + + +@skip_if_not_sqlite +def test_read_dataset_remote_update_flag( + studio_token, + test_session, + remote_dataset_multi_version, + mock_dataset_info_endpoint, + mock_export_endpoint_with_urls, + mock_export_status_completed, + mock_s3_parquet_download, + mock_dataset_rows_fetcher_status_check, + requests_mock, +): + """Test read_dataset with update=True flag to force remote check.""" + + # Mock the Studio API responses + mock_dataset_info_endpoint(remote_dataset_multi_version) + mock_s3_parquet_download() + + # First read - downloads version 1.0.0 + ds1 = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version="1.0.0", + session=test_session, + ) + assert dc.datasets().to_values("version") == ["1.0.0"] + assert ds1.to_values("version")[0] == "1.0.0" + + # Second read with update=True with the exact version + # returns the same + ds2 = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version="1.0.0", + update=True, + session=test_session, + ) + assert dc.datasets().to_values("version") == ["1.0.0"] + assert ds2.to_values("version")[0] == "1.0.0" + + # Third read with update=False even with version specifier + # that allows for newer version still bring the same version + # as the one already downloaded + ds3 = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version=">=1.0.0", + update=False, + session=test_session, + ) + assert dc.datasets().to_values("version") == ["1.0.0"] + assert ds3.to_values("version")[0] == "1.0.0" + + # Finally, read with update=False even with version specifier + # that allows for newer version still bring the same version + # as the one already downloaded + ds4 = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version=">=1.0.0", + update=True, + session=test_session, + ) + + assert ds4.to_values("version")[0] == "2.0.0" + assert dc.datasets().to_values("version") == ["1.0.0", "2.0.0"] + + +@skip_if_not_sqlite +def test_read_dataset_remote_version_specifiers( + studio_token, + test_session, + remote_dataset_multi_version, + mock_dataset_info_endpoint, + mock_export_endpoint_with_urls, + mock_export_status_completed, + mock_s3_parquet_download, + mock_dataset_rows_fetcher_status_check, +): + """Test read_dataset with version specifiers on remote datasets.""" + + # Mock the Studio API responses + mock_dataset_info_endpoint(remote_dataset_multi_version) + mock_s3_parquet_download() + + # Test reading with version specifier ">=1.0.0" should get latest (2.0.0) + ds = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version=">=1.0.0", + session=test_session, + ) + + # Should get the latest version matching the specifier (2.0.0) + assert ds.dataset.name == "dogs" + dataset_version = ds.dataset.get_version("2.0.0") + assert dataset_version is not None + assert dataset_version.uuid == REMOTE_DATASET_UUID_V2 + assert dc.datasets().to_values("version") == ["2.0.0"] + assert ds.to_values("version")[0] == "2.0.0" + + +@skip_if_not_sqlite +def test_read_dataset_remote_version_specifier_no_match( + studio_token, + test_session, + remote_dataset_multi_version, + mock_dataset_info_endpoint, + mock_dataset_rows_fetcher_status_check, +): + """Test read_dataset with version specifier that doesn't match.""" + + mock_dataset_info_endpoint(remote_dataset_multi_version) + + # Test version specifier that doesn't match any existing version + with pytest.raises(DatasetVersionNotFoundError) as exc_info: + dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version=">=3.0.0", + session=test_session, + ) + + assert "No dataset" in str(exc_info.value) + assert "version matching specifier >=3.0.0" in str(exc_info.value) + + +@skip_if_not_sqlite +def test_read_dataset_remote_not_found( + studio_token, + test_session, + mock_dataset_info_not_found, +): + """Test read_dataset when remote dataset is not found.""" + + with pytest.raises(DatasetNotFoundError) as exc_info: + dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.nonexistent", + version="1.0.0", + session=test_session, + ) + + expected_msg = ( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.nonexistent not found" + ) + assert expected_msg in str(exc_info.value) + + +@skip_if_not_sqlite +def test_read_dataset_remote_version_not_found( + studio_token, + test_session, + remote_dataset_single_version, + mock_dataset_info_endpoint, + mock_dataset_rows_fetcher_status_check, +): + """Test read_dataset when requested version doesn't exist on remote.""" + + mock_dataset_info_endpoint(remote_dataset_single_version) + + with pytest.raises(DataChainError) as exc_info: + dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version="5.0.0", + session=test_session, + ) + + assert "Dataset dogs doesn't have version 5.0.0 on server" in str(exc_info.value) + + +@skip_if_not_sqlite +def test_read_dataset_remote_latest_version( + studio_token, + test_session, + remote_dataset_multi_version, + mock_dataset_info_endpoint, + mock_export_endpoint_with_urls, + mock_export_status_completed, + mock_s3_parquet_download, + mock_dataset_rows_fetcher_status_check, +): + """Test read_dataset without version parameter should get latest version.""" + + # Mock the Studio API responses + mock_dataset_info_endpoint(remote_dataset_multi_version) + mock_s3_parquet_download() + + # Read without specifying version should get latest + ds = dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + session=test_session, + ) + + # Should get the latest version (2.0.0) + assert ds.dataset.name == "dogs" + dataset_version = ds.dataset.get_version("2.0.0") + assert dataset_version is not None + assert dataset_version.uuid == REMOTE_DATASET_UUID_V2 + assert dc.datasets().to_values("version") == ["2.0.0"] + assert ds.to_values("version")[0] == "2.0.0" + + +@skip_if_not_sqlite +def test_read_dataset_remote_export_failed( + studio_token, + test_session, + remote_dataset_single_version, + mock_dataset_info_endpoint, + mock_export_endpoint_with_id, + mock_export_status_failed, + mock_dataset_rows_fetcher_status_check, +): + """Test read_dataset when remote dataset export fails.""" + + mock_dataset_info_endpoint(remote_dataset_single_version) + + with pytest.raises(DataChainError) as exc_info: + dc.read_dataset( + f"{REMOTE_NAMESPACE_NAME}.{REMOTE_PROJECT_NAME}.dogs", + version="1.0.0", + session=test_session, + ) + + assert "Dataset export failed in Studio" in str(exc_info.value) diff --git a/tests/func/test_read_dataset_version_specifiers.py b/tests/func/test_read_dataset_version_specifiers.py new file mode 100644 index 000000000..6fc686265 --- /dev/null +++ b/tests/func/test_read_dataset_version_specifiers.py @@ -0,0 +1,88 @@ +""" +Functional tests for read_dataset with PEP 440 version specifiers. +""" + +import pytest + +import datachain as dc +from datachain.error import DatasetVersionNotFoundError + + +def test_read_dataset_version_specifiers(test_session): + """Test read_dataset with various PEP 440 version specifiers.""" + # Create a dataset with multiple versions + dataset_name = "test_version_specifiers" + + for version in ["1.0.0", "1.1.0", "1.2.0", "2.0.0"]: + ( + dc.read_values(data=[1, 2], session=test_session) + .mutate(dataset_version=version) + .save(dataset_name, version=version) + ) + + # Test exact version specifier + result = dc.read_dataset(dataset_name, version="==1.1.0", session=test_session) + assert result.to_values("dataset_version")[0] == "1.1.0" + + # Test greater than or equal specifier - should get latest (2.0.0) + result = dc.read_dataset(dataset_name, version=">=1.1.0", session=test_session) + assert result.to_values("dataset_version")[0] == "2.0.0" + + # Test less than specifier - should get 1.2.0 (latest before 2.0.0) + result = dc.read_dataset(dataset_name, version="<2.0.0", session=test_session) + assert result.to_values("dataset_version")[0] == "1.2.0" + + # Test compatible release specifier - should get latest 1.x (1.2.0) + result = dc.read_dataset(dataset_name, version="~=1.0", session=test_session) + assert result.to_values("dataset_version")[0] == "1.2.0" + + # Test version pattern - should get latest 1.x (1.2.0) + result = dc.read_dataset(dataset_name, version="==1.*", session=test_session) + assert result.to_values("dataset_version")[0] == "1.2.0" + + # Test complex specifier - should get 1.2.0 + result = dc.read_dataset( + dataset_name, version=">=1.1.0,<2.0.0", session=test_session + ) + assert result.to_values("dataset_version")[0] == "1.2.0" + + +def test_read_dataset_version_specifiers_no_match(test_session): + """Test read_dataset with version specifiers that don't match any version.""" + # Create a dataset with a single version + dataset_name = "test_no_match_specifiers" + + ( + dc.read_values(data=[1, 2], session=test_session) + .mutate(dataset_version="1.0.0") + .save(dataset_name, version="1.0.0") + ) + + # Test version specifier that doesn't match any existing version + with pytest.raises(DatasetVersionNotFoundError) as exc_info: + dc.read_dataset(dataset_name, version=">=2.0.0", session=test_session) + + assert ( + "No dataset test_no_match_specifiers version matching specifier >=2.0.0" + in str(exc_info.value) + ) + + +def test_read_dataset_version_specifiers_exact_version(test_session): + """Test that version specifiers work alongside with exact version reads.""" + # Create a dataset with multiple versions + dataset_name = "test_backward_compatibility" + + ( + dc.read_values(data=[1, 2], session=test_session) + .mutate(dataset_version="1.0.0") + .save(dataset_name, version="1.0.0") + ) + + # Test reading by exact version + result = dc.read_dataset(dataset_name, version="1.0.0", session=test_session) + assert result.to_values("dataset_version")[0] == "1.0.0" + + # Test reading by exact version int - backward compatibility + result = dc.read_dataset(dataset_name, version=1, session=test_session) + assert result.to_values("dataset_version")[0] == "1.0.0" diff --git a/tests/func/test_retry.py b/tests/func/test_retry.py index cce7e3dcf..c8accd2c9 100644 --- a/tests/func/test_retry.py +++ b/tests/func/test_retry.py @@ -5,7 +5,7 @@ import datachain as dc from datachain import C, DataModel -from datachain.error import DatasetVersionNotFoundError +from datachain.error import DatasetNotFoundError if TYPE_CHECKING: from datachain import DataChain @@ -173,7 +173,7 @@ def successful_process(id: int, content: str) -> ProcessingResult: ) # Should not create version 1.0.1 since no retry was needed - with pytest.raises(DatasetVersionNotFoundError): + with pytest.raises(DatasetNotFoundError): dc.read_dataset("successful_data", version="1.0.1", session=test_session) diff --git a/tests/unit/lib/test_datachain.py b/tests/unit/lib/test_datachain.py index 56b101d12..9a1efafcc 100644 --- a/tests/unit/lib/test_datachain.py +++ b/tests/unit/lib/test_datachain.py @@ -3253,10 +3253,19 @@ def test_delete_dataset_from_studio_not_found( assert str(exc_info.value) == error_message -def test_delete_dataset_cached_from_studio(test_session, project): +def test_delete_dataset_cached_from_studio( + test_session, project, studio_token, requests_mock +): ds_full_name = f"{project.namespace.name}.{project.name}.fibonacci" dc.read_values(fib=[1, 1, 2, 3, 5, 8], session=test_session).save(ds_full_name) + error_message = f"Dataset {ds_full_name} not found" + requests_mock.get( + f"{STUDIO_URL}/api/datachain/datasets/info", + json={"message": error_message}, + status_code=404, + ) + dc.delete_dataset(ds_full_name) with pytest.raises(DatasetNotFoundError):