diff --git a/dagshub/data_engine/client/data_client.py b/dagshub/data_engine/client/data_client.py index fa8bacb1..32fec418 100644 --- a/dagshub/data_engine/client/data_client.py +++ b/dagshub/data_engine/client/data_client.py @@ -83,6 +83,14 @@ def create_datasource(self, ds: "DatasourceState") -> DatasourceResult: res = self._exec(q, params) return dacite.from_dict(DatasourceResult, res["createDatasource"], config=dacite_config) + def copy_datasource( + self, source: Union[int, str], name: str, query_input: Optional[Dict[str, Any]] = None + ) -> DatasourceResult: + q = GqlMutations.copy_datasource() + params = GqlMutations.copy_datasource_params(source=source, name=name, query=query_input) + res = self._exec(q, params) + return dacite.from_dict(DatasourceResult, res["copyDatasource"], config=dacite_config) + def head(self, datasource: "Datasource", size: Optional[int] = None) -> QueryResult: """ Retrieve a subset of data from the datasource headers. diff --git a/dagshub/data_engine/client/gql_mutations.py b/dagshub/data_engine/client/gql_mutations.py index 1b53f9fd..6f784777 100644 --- a/dagshub/data_engine/client/gql_mutations.py +++ b/dagshub/data_engine/client/gql_mutations.py @@ -27,6 +27,36 @@ def create_datasource(): def create_datasource_params(name: str, url: str, ds_type: DatasourceType): return {"name": name, "url": url, "dsType": str(ds_type.value)} + @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", + "type", + "origin {sourceDatasourceId sourceRepoId sourceName sourceRootUrl sourceType " + "sourceBackingType creatorId createdAt query}", + ] + ) + .param_validator(Validators.query_input_validator("query")) + ) + + @staticmethod + def copy_datasource_params(source: Union[int, str], name: str, query: Optional[Dict[str, Any]]): + return {"source": source, "name": name, "query": query} + @staticmethod @functools.lru_cache() def update_metadata(): diff --git a/dagshub/data_engine/client/models.py b/dagshub/data_engine/client/models.py index 9e0f2427..4cee9cfb 100644 --- a/dagshub/data_engine/client/models.py +++ b/dagshub/data_engine/client/models.py @@ -101,6 +101,19 @@ 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 + sourceBackingType: str + creatorId: Union[str, int] + createdAt: int + query: str + + @dataclass class DatasourceResult: id: Union[str, int] @@ -110,6 +123,7 @@ class DatasourceResult: preprocessingStatus: PreprocessingStatus type: DatasourceType metadataFields: Optional[List[MetadataFieldSchema]] + origin: Optional[DatasourceOriginResult] = None @dataclass diff --git a/dagshub/data_engine/datasources.py b/dagshub/data_engine/datasources.py index e2c3be75..918f739e 100644 --- a/dagshub/data_engine/datasources.py +++ b/dagshub/data_engine/datasources.py @@ -60,6 +60,35 @@ def create(*args, **kwargs) -> Datasource: return create_datasource(*args, **kwargs) +def copy_datasource( + repo: str, + source: Union[str, int, Datasource], + name: str, + query: Optional[Dict] = None, +) -> Datasource: + """Materialize a new datasource from a query over an existing datasource. + + The copy runs asynchronously. Its preprocessing status exposes progress through the + same APIs used for newly scanned datasources. When ``source`` is a Datasource, its + current filters, projection, limit, and as-of timestamp are used automatically. + Version history before the materialized state is not copied. + """ + if isinstance(source, Datasource): + source_id = source.source.id + if query is None: + query = source.serialize_gql_query_input() + else: + source_id = source + if source_id is None: + raise ValueError("source datasource must have an id") + result = DataClient(repo).copy_datasource( + source=source_id, + name=name, + query_input=query, + ) + return Datasource(DatasourceState.from_gql_result(repo, result)) + + def get_or_create(repo: str, name: str, path: str, revision: Optional[str] = None) -> Datasource: """ First attempts to get the repo datasource with the given name, and only if that fails, @@ -236,6 +265,7 @@ def _load_datasources_from_run( __all__ = [ create_datasource.__name__, create.__name__, + copy_datasource.__name__, create_from_bucket.__name__, create_from_repo.__name__, get_datasource.__name__, diff --git a/dagshub/data_engine/model/datasource.py b/dagshub/data_engine/model/datasource.py index bbeab214..2a0fd51f 100644 --- a/dagshub/data_engine/model/datasource.py +++ b/dagshub/data_engine/model/datasource.py @@ -173,6 +173,15 @@ def __init__( self.ngrok_listener = None + def copy_to( + self, + name: str, + ) -> "Datasource": + """Asynchronously materialize this datasource's current query as a datasource.""" + from dagshub.data_engine.datasources import copy_datasource + + return copy_datasource(self.source.repo, self, name) + @property def has_explicit_context(self): return self._explicit_update_ctx is not None diff --git a/dagshub/data_engine/model/datasource_state.py b/dagshub/data_engine/model/datasource_state.py index 93f50e08..efd5105c 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 ( + DatasourceType, + DatasourceResult, + DatasourceOriginResult, + 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) @@ -209,6 +216,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_copy.py b/tests/data_engine/test_datasource_copy.py new file mode 100644 index 00000000..6393c293 --- /dev/null +++ b/tests/data_engine/test_datasource_copy.py @@ -0,0 +1,84 @@ +from dagshub.data_engine.client.gql_mutations import GqlMutations +from dagshub.data_engine.client.data_client import DataClient +from dagshub.data_engine.client.models import PreprocessingStatus +from dagshub.data_engine.model.datasource import Datasource +from dagshub.data_engine.model.datasource_state import DatasourceState + + +def test_copy_datasource_mutation_and_params(): + query = GqlMutations.copy_datasource().generate() + + assert "mutation copyDatasource" in query + assert "$source: ID!" in query + assert "$query: QueryInput" in query + assert "copyDatasource(source: $source, name: $name, query: $query)" in query + query_input = {"select": [{"name": "score", "alias": "confidence"}], "limit": 20} + assert GqlMutations.copy_datasource_params(12, "snapshot", query_input) == { + "source": 12, + "name": "snapshot", + "query": query_input, + } + + +def test_copy_datasource_params_accept_latest_state(): + assert GqlMutations.copy_datasource_params("12", "latest", None)["query"] is None + + +def test_data_client_returns_copied_datasource(monkeypatch): + client = object.__new__(DataClient) + captured = {} + + def fake_exec(query, params): + captured.update(params) + return { + "copyDatasource": { + "id": "13", + "name": "snapshot", + "rootUrl": "repo://owner/repo/main:data", + "integrationStatus": "VALID", + "preprocessingStatus": "IN_PROGRESS", + "type": "REPOSITORY", + "origin": { + "sourceDatasourceId": "12", + "sourceRepoId": "7", + "sourceName": "source", + "sourceRootUrl": "repo://owner/repo/main:data", + "sourceType": "REPOSITORY", + "sourceBackingType": "postgres", + "creatorId": "3", + "createdAt": 1_700_000_001, + "query": '{"asOf":1700000000,"limit":20}', + }, + } + } + + monkeypatch.setattr(client, "_exec", fake_exec) + query_input = {"asOf": 1_700_000_000, "limit": 20} + result = client.copy_datasource(12, "snapshot", query_input) + + assert captured == {"source": 12, "name": "snapshot", "query": query_input} + assert result.id == "13" + assert result.preprocessingStatus == PreprocessingStatus.IN_PROGRESS + assert result.origin is not None + assert result.origin.sourceName == "source" + assert result.origin.sourceBackingType == "postgres" + + +def test_datasource_copy_materializes_current_query(monkeypatch): + state = object.__new__(DatasourceState) + state.repo = "owner/repo" + state.id = 12 + queried = Datasource(state).select("score").limit(5) + captured = {} + + def fake_copy(repo, source, name): + captured.update(repo=repo, source=source, name=name, query=source.serialize_gql_query_input()) + return "copied" + + monkeypatch.setattr("dagshub.data_engine.datasources.copy_datasource", fake_copy) + + assert queried.copy_to("projection") == "copied" + assert captured["repo"] == "owner/repo" + assert captured["source"] is queried + assert captured["name"] == "projection" + assert captured["query"] == {"select": [{"name": "score"}], "query": None, "limit": 5}