From 0278a095fc9809d4f5f9db12ba633aece1116d7e Mon Sep 17 00:00:00 2001 From: Guy Smoilovsky Date: Fri, 28 Aug 2026 11:06:35 +0300 Subject: [PATCH 1/2] feat(data-engine): add datasource copy API --- dagshub/data_engine/client/data_client.py | 12 ++++++ dagshub/data_engine/client/gql_mutations.py | 32 +++++++++++++++ dagshub/data_engine/client/models.py | 15 ++++++- dagshub/data_engine/datasources.py | 14 +++++++ dagshub/data_engine/model/datasource.py | 15 +++++++ dagshub/data_engine/model/datasource_state.py | 10 ++++- tests/data_engine/test_datasource.py | 39 ++++++++++++++++++- 7 files changed, 134 insertions(+), 3 deletions(-) diff --git a/dagshub/data_engine/client/data_client.py b/dagshub/data_engine/client/data_client.py index fa8bacb17..19b10ca23 100644 --- a/dagshub/data_engine/client/data_client.py +++ b/dagshub/data_engine/client/data_client.py @@ -329,6 +329,18 @@ def delete_datasource(self, datasource: "Datasource"): params = GqlMutations.delete_datasource_params(datasource_id=datasource.source.id) return self._exec(q, params) + def copy_datasource(self, datasource: "Datasource", name: str) -> DatasourceResult: + """Create a new datasource from the source's current query snapshot.""" + assert datasource.source.id is not None + q = GqlMutations.copy_datasource() + params = GqlMutations.copy_datasource_params( + source_id=datasource.source.id, + name=name, + query=datasource.serialize_gql_query_input(), + ) + res = self._exec(q, params)["copyDatasource"] + return dacite.from_dict(DatasourceResult, res, config=dacite_config) + def scan_datasource(self, datasource: "Datasource", options: Optional[List[ScanOption]]): """ Initiate a scan operation on the specified datasource. diff --git a/dagshub/data_engine/client/gql_mutations.py b/dagshub/data_engine/client/gql_mutations.py index 1b53f9fde..b8ba316be 100644 --- a/dagshub/data_engine/client/gql_mutations.py +++ b/dagshub/data_engine/client/gql_mutations.py @@ -168,6 +168,38 @@ def delete_datasource_params(datasource_id: Union[int, str]): "id": datasource_id, } + @staticmethod + @functools.lru_cache() + def copy_datasource(): + return ( + GqlQuery() + .operation( + "mutation", + name="copyDatasource", + input={"$source": "ID!", "$name": "String!", "$query": "QueryInput"}, + ) + .query( + "copyDatasource", + input={"source": "$source", "name": "$name", "query": "$query"}, + ) + .fields( + [ + "id", + "name", + "rootUrl", + "integrationStatus", + "preprocessingStatus", + "metadataFields {name valueType multiple tags}", + "type", + "origin {sourceDatasourceId sourceRepoId sourceName sourceRootUrl sourceType creatorId createdAt query}", + ] + ) + ) + + @staticmethod + def copy_datasource_params(source_id: Union[int, str], name: str, query: Dict[str, Any]): + return {"source": source_id, "name": name, "query": query} + @staticmethod @functools.lru_cache() def scan_datasource(): diff --git a/dagshub/data_engine/client/models.py b/dagshub/data_engine/client/models.py index 9e0f24271..39e5ad75b 100644 --- a/dagshub/data_engine/client/models.py +++ b/dagshub/data_engine/client/models.py @@ -101,6 +101,18 @@ def from_metadata_field_schema(mfs: MetadataFieldSchema) -> "MetadataSelectField ) +@dataclass +class DatasourceOriginResult: + sourceDatasourceId: Union[str, int] + sourceRepoId: Union[str, int] + sourceName: str + sourceRootUrl: str + sourceType: DatasourceType + creatorId: Union[str, int] + createdAt: datetime.datetime + query: str + + @dataclass class DatasourceResult: id: Union[str, int] @@ -109,7 +121,8 @@ class DatasourceResult: integrationStatus: IntegrationStatus preprocessingStatus: PreprocessingStatus type: DatasourceType - metadataFields: Optional[List[MetadataFieldSchema]] + metadataFields: Optional[List[MetadataFieldSchema]] = None + origin: Optional[DatasourceOriginResult] = None @dataclass diff --git a/dagshub/data_engine/datasources.py b/dagshub/data_engine/datasources.py index e2c3be75b..9b798a413 100644 --- a/dagshub/data_engine/datasources.py +++ b/dagshub/data_engine/datasources.py @@ -168,6 +168,19 @@ def get_datasources(repo: str) -> List[Datasource]: return [Datasource(DatasourceState.from_gql_result(repo, source)) for source in sources] +def copy_datasource(source: Datasource, name: str) -> Datasource: + """Create a datasource containing a snapshot of ``source``'s current query. + + The copy runs asynchronously on DagsHub. Call ``wait_until_ready`` on the + returned datasource before reading it. + + Args: + source: Datasource, including any filter/select query, to copy. + name: Name of the new datasource. + """ + return source.copy_datasource(name) + + def get_from_mlflow( run: Optional[Union["mlflow.entities.Run", str]] = None, artifact_name: Optional[str] = None ) -> Dict[str, Datasource]: @@ -238,6 +251,7 @@ def _load_datasources_from_run( create.__name__, create_from_bucket.__name__, create_from_repo.__name__, + copy_datasource.__name__, get_datasource.__name__, get_datasources.__name__, get.__name__, diff --git a/dagshub/data_engine/model/datasource.py b/dagshub/data_engine/model/datasource.py index bbeab214e..c3bfd8625 100644 --- a/dagshub/data_engine/model/datasource.py +++ b/dagshub/data_engine/model/datasource.py @@ -702,6 +702,21 @@ def delete_source(self, force: bool = False): return self.source.client.delete_datasource(self) + def copy_datasource(self, name: str) -> "Datasource": + """Create a datasource containing a snapshot of this datasource's current query. + + The copy runs asynchronously on DagsHub. Use :meth:`wait_until_ready` on + the returned datasource before reading it. + + Args: + name: Name of the new datasource. + + Returns: + The newly created datasource. + """ + result = self.source.client.copy_datasource(self, name) + return Datasource(DatasourceState.from_gql_result(self.source.repo, result)) + def delete_dataset(self, force: bool = False): """ Deletes the dataset, if this object was created from a dataset diff --git a/dagshub/data_engine/model/datasource_state.py b/dagshub/data_engine/model/datasource_state.py index ff2809955..8f144d45d 100644 --- a/dagshub/data_engine/model/datasource_state.py +++ b/dagshub/data_engine/model/datasource_state.py @@ -7,7 +7,13 @@ from os import PathLike from dagshub.common.api.repo import RepoAPI, PathNotFoundError from dagshub.data_engine.client.data_client import DataClient -from dagshub.data_engine.client.models import DatasourceType, DatasourceResult, PreprocessingStatus, MetadataFieldSchema +from dagshub.data_engine.client.models import ( + DatasourceOriginResult, + DatasourceType, + DatasourceResult, + PreprocessingStatus, + MetadataFieldSchema, +) from dagshub.data_engine.model.datapoint import Datapoint from dagshub.data_engine.model.errors import DatasourceAlreadyExistsError, DatasourceNotFoundError from dagshub.common.util import multi_urljoin @@ -43,6 +49,7 @@ class DatasourceState: client: DataClient = field(init=False) repoApi: RepoAPI = field(init=False) metadata_fields: List[MetadataFieldSchema] = field(init=False) + origin: Optional[DatasourceOriginResult] = field(init=False, default=None) _revision: Optional[str] = field(init=False, default=None) @@ -215,6 +222,7 @@ def _update_from_ds_result(self, ds: DatasourceResult): self.source_type = ds.type self.preprocessing_status = ds.preprocessingStatus self.metadata_fields = [] if ds.metadataFields is None else ds.metadataFields + self.origin = ds.origin if self.source_type == DatasourceType.REPOSITORY: self.revision = self.path_parts()["revision"] diff --git a/tests/data_engine/test_datasource.py b/tests/data_engine/test_datasource.py index e6f6e0dce..67402a897 100644 --- a/tests/data_engine/test_datasource.py +++ b/tests/data_engine/test_datasource.py @@ -11,7 +11,14 @@ import dagshub.common.config from dagshub.common.util import wrap_bytes from dagshub.data_engine.annotation import MetadataAnnotations -from dagshub.data_engine.client.models import MetadataFieldSchema +from dagshub.data_engine.client.models import ( + DatasourceOriginResult, + DatasourceResult, + DatasourceType, + IntegrationStatus, + MetadataFieldSchema, + PreprocessingStatus, +) from dagshub.data_engine.dtypes import MetadataFieldType, ReservedTags from dagshub.data_engine.model.datapoint import Datapoint from dagshub.data_engine.model.datasource import DatapointMetadataUpdateEntry, Datasource, MetadataContextManager @@ -54,6 +61,36 @@ def test_default_behavior(ds, metadata_df): assert expected == actual +def test_copy_datasource_returns_new_source_with_origin(ds): + origin = DatasourceOriginResult( + sourceDatasourceId=ds.source.id, + sourceRepoId=12, + sourceName=ds.source.name, + sourceRootUrl=ds.source.path, + sourceType=DatasourceType.REPOSITORY, + creatorId=7, + createdAt=datetime.datetime.now(datetime.timezone.utc), + query="{}", + ) + ds.source.client.copy_datasource.return_value = DatasourceResult( + id=2, + name="copied-datasource", + rootUrl=ds.source.path, + integrationStatus=IntegrationStatus.VALID, + preprocessingStatus=PreprocessingStatus.IN_PROGRESS, + type=DatasourceType.REPOSITORY, + origin=origin, + ) + + copied = ds.copy_datasource("copied-datasource") + + ds.source.client.copy_datasource.assert_called_once_with(ds, "copied-datasource") + assert copied.source.id == 2 + assert copied.source.name == "copied-datasource" + assert copied.source.preprocessing_status == PreprocessingStatus.IN_PROGRESS + assert copied.source.origin == origin + + @pytest.mark.parametrize("column", ["key3", 3]) def test_column_arg(ds, metadata_df, column): actual = Datasource._df_to_metadata(ds, metadata_df, column) From af06b68d20255bc53bf74202155ca00dc650b0dc Mon Sep 17 00:00:00 2001 From: Guy Smoilovsky Date: Tue, 1 Sep 2026 11:49:11 +0300 Subject: [PATCH 2/2] Preserve datasource copy origin --- dagshub/data_engine/client/gql_mutations.py | 3 +- dagshub/data_engine/model/datasource_state.py | 3 +- tests/data_engine/test_datasource.py | 40 +++++++++++++++++++ 3 files changed, 44 insertions(+), 2 deletions(-) diff --git a/dagshub/data_engine/client/gql_mutations.py b/dagshub/data_engine/client/gql_mutations.py index b8ba316be..7bb089748 100644 --- a/dagshub/data_engine/client/gql_mutations.py +++ b/dagshub/data_engine/client/gql_mutations.py @@ -191,7 +191,8 @@ def copy_datasource(): "preprocessingStatus", "metadataFields {name valueType multiple tags}", "type", - "origin {sourceDatasourceId sourceRepoId sourceName sourceRootUrl sourceType creatorId createdAt query}", + "origin {sourceDatasourceId sourceRepoId sourceName sourceRootUrl sourceType " + "creatorId createdAt query}", ] ) ) diff --git a/dagshub/data_engine/model/datasource_state.py b/dagshub/data_engine/model/datasource_state.py index 8f144d45d..7e687262e 100644 --- a/dagshub/data_engine/model/datasource_state.py +++ b/dagshub/data_engine/model/datasource_state.py @@ -222,7 +222,8 @@ def _update_from_ds_result(self, ds: DatasourceResult): self.source_type = ds.type self.preprocessing_status = ds.preprocessingStatus self.metadata_fields = [] if ds.metadataFields is None else ds.metadataFields - self.origin = ds.origin + if ds.origin is not None: + self.origin = ds.origin if self.source_type == DatasourceType.REPOSITORY: self.revision = self.path_parts()["revision"] diff --git a/tests/data_engine/test_datasource.py b/tests/data_engine/test_datasource.py index 67402a897..73810512e 100644 --- a/tests/data_engine/test_datasource.py +++ b/tests/data_engine/test_datasource.py @@ -22,6 +22,7 @@ from dagshub.data_engine.dtypes import MetadataFieldType, ReservedTags from dagshub.data_engine.model.datapoint import Datapoint from dagshub.data_engine.model.datasource import DatapointMetadataUpdateEntry, Datasource, MetadataContextManager +from dagshub.data_engine.model.datasource_state import DatasourceState from dagshub.data_engine.model.errors import DataEngineGqlError from dagshub.data_engine.model.metadata import MultipleDataTypesUploadedError, StringFieldValueTooLongError from dagshub.data_engine.model.query_result import QueryResult @@ -91,6 +92,45 @@ def test_copy_datasource_returns_new_source_with_origin(ds): assert copied.source.origin == origin +def test_copy_datasource_preserves_origin_when_refresh_omits_it(ds): + origin = DatasourceOriginResult( + sourceDatasourceId=ds.source.id, + sourceRepoId=12, + sourceName=ds.source.name, + sourceRootUrl=ds.source.path, + sourceType=DatasourceType.REPOSITORY, + creatorId=7, + createdAt=datetime.datetime.now(datetime.timezone.utc), + query="{}", + ) + copied = DatasourceState.from_gql_result( + ds.source.repo, + DatasourceResult( + id=2, + name="copied-datasource", + rootUrl=ds.source.path, + integrationStatus=IntegrationStatus.VALID, + preprocessingStatus=PreprocessingStatus.IN_PROGRESS, + type=DatasourceType.REPOSITORY, + origin=origin, + ), + ) + + copied._update_from_ds_result( + DatasourceResult( + id=2, + name="copied-datasource", + rootUrl=ds.source.path, + integrationStatus=IntegrationStatus.VALID, + preprocessingStatus=PreprocessingStatus.READY, + type=DatasourceType.REPOSITORY, + ) + ) + + assert copied.preprocessing_status == PreprocessingStatus.READY + assert copied.origin == origin + + @pytest.mark.parametrize("column", ["key3", 3]) def test_column_arg(ds, metadata_df, column): actual = Datasource._df_to_metadata(ds, metadata_df, column)