Skip to content
Draft
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
8 changes: 8 additions & 0 deletions dagshub/data_engine/client/data_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
30 changes: 30 additions & 0 deletions dagshub/data_engine/client/gql_mutations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down
14 changes: 14 additions & 0 deletions dagshub/data_engine/client/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -110,6 +123,7 @@ class DatasourceResult:
preprocessingStatus: PreprocessingStatus
type: DatasourceType
metadataFields: Optional[List[MetadataFieldSchema]]
origin: Optional[DatasourceOriginResult] = None


@dataclass
Expand Down
30 changes: 30 additions & 0 deletions dagshub/data_engine/datasources.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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__,
Expand Down
9 changes: 9 additions & 0 deletions dagshub/data_engine/model/datasource.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
10 changes: 9 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 (
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
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 @@ -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"]

Expand Down
84 changes: 84 additions & 0 deletions tests/data_engine/test_datasource_copy.py
Original file line number Diff line number Diff line change
@@ -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}
Loading