Skip to content
Merged
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
12 changes: 12 additions & 0 deletions dagshub/data_engine/client/data_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
33 changes: 33 additions & 0 deletions dagshub/data_engine/client/gql_mutations.py
Original file line number Diff line number Diff line change
Expand Up @@ -168,6 +168,39 @@ 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():
Expand Down
15 changes: 14 additions & 1 deletion dagshub/data_engine/client/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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
Expand Down
14 changes: 14 additions & 0 deletions dagshub/data_engine/datasources.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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__,
Expand Down
15 changes: 15 additions & 0 deletions dagshub/data_engine/model/datasource.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Comment thread
coderabbitai[bot] marked this conversation as resolved.

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
Expand Down
11 changes: 10 additions & 1 deletion dagshub/data_engine/model/datasource_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -215,6 +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
if ds.origin is not None:
self.origin = ds.origin
if self.source_type == DatasourceType.REPOSITORY:
self.revision = self.path_parts()["revision"]

Expand Down
79 changes: 78 additions & 1 deletion tests/data_engine/test_datasource.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,10 +11,18 @@
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
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
Expand Down Expand Up @@ -54,6 +62,75 @@ 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


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)
Expand Down
Loading