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
37 changes: 16 additions & 21 deletions dreadnode/api/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -781,6 +781,22 @@ def get_user_data_credentials(
response = self.request("GET", "/user-data/credentials", params=params)
return UserDataCredentials(**response.json())

def get_dataset_access_credentials(
self,
dataset_id: UUID,
) -> UserDataCredentials:
"""
Retrieves dataset access credentials for a specific dataset.

Args:
dataset_id (UUID): The dataset identifier.
Returns:
The user data credentials object.
"""
params: dict[str, str] = {"dataset_id": str(dataset_id)}
response = self.request("GET", f"/datasets/{dataset_id!s}/credentials", params=params)
return UserDataCredentials(**response.json())

# Container registry access

def get_platform_registry_credentials(self) -> ContainerRegistryCredentials:
Expand Down Expand Up @@ -991,27 +1007,6 @@ def get_dataset(
response = self.request("GET", f"/datasets/{dataset_id_or_key}")
return DatasetMetadata(**response.json())

def update_dataset(
self,
dataset_id_or_key: str | UUID,
dataset: CreateDatasetRequest,
) -> DatasetMetadata:
"""
Updates an existing dataset.

Args:
dataset_id_or_key (str | UUID): The dataset identifier.
dataset (DatasetCreateRequest): The dataset update request object.

Returns:
DatasetMetadata: The updated DatasetMetadata object.
"""

payload: dict[str, t.Any] = dataset.model_dump()

response = self.request("PUT", f"/datasets/{dataset_id_or_key}", json_data=payload)
return DatasetMetadata(**response.json())

def delete_dataset(self, dataset_id_or_key: str | UUID) -> None:
"""
Deletes a specific dataset.
Expand Down
11 changes: 11 additions & 0 deletions dreadnode/api/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,8 @@
)
from ulid import ULID

from dreadnode.util import valid_version

AnyDict = dict[str, t.Any]


Expand All @@ -29,6 +31,12 @@ def _validate_key(key: str) -> str:
return key


def _validate_version(version: str) -> str:
if not valid_version(version):
raise ValidationError("Version must follow semantic versioning (e.g., '1.0.0').")
return version


# User


Expand Down Expand Up @@ -590,6 +598,9 @@ class CreateDatasetRequest(BaseModel):
"""Unique identifier for the organization owning the dataset."""
key: t.Annotated[str, BeforeValidator(_validate_key)]
"""Unique identifier of the dataset."""
version: t.Annotated[str, BeforeValidator(_validate_version)]
tags: list[str] | None = None
"""Optional list of tags associated with the dataset."""


class CreateDatasetResponse(BaseModel):
Expand Down
5 changes: 5 additions & 0 deletions dreadnode/common_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,3 +69,8 @@ def __lt__(self, __other: te.Self) -> bool: ... # noqa: PYI063

class SupportsLe(t.Protocol):
def __le__(self, __other: te.Self) -> bool: ... # noqa: PYI063


# Versioning

VersionStrategy = t.Literal["major", "minor", "patch"]
5 changes: 2 additions & 3 deletions dreadnode/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,7 +93,7 @@ def set_size_and_row_count(self) -> None:

try:
self.row_count = self.ds.count_rows()
except Exception:
except pa.ArrowException:
self.row_count = 0

def save_metadata(self, path: str, fs: FileSystem) -> None:
Expand Down Expand Up @@ -338,5 +338,4 @@ def load_dataset(
return Dataset(ds=dataset, metadata=metadata, manifest=manifest, materialize=materialize)
except Exception as e:
print_info(f"[!] Failed to load dataset from remote: {e}")

raise FileNotFoundError(f"[!] Dataset not found: {uri}")
raise FileNotFoundError(f"[!] Dataset not found: {uri}") from e
45 changes: 45 additions & 0 deletions dreadnode/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@
from opentelemetry import propagate
from opentelemetry.exporter.otlp.proto.http import Compression
from opentelemetry.sdk.trace.export import BatchSpanProcessor
from packaging.version import InvalidVersion
from packaging.version import parse as parse_version

from dreadnode import dataset
from dreadnode.api.client import ApiClient
Expand All @@ -35,6 +37,7 @@
AnyDict,
Inherited,
JsonValue,
VersionStrategy,
)
from dreadnode.constants import (
DEFAULT_LOCAL_STORAGE_DIR,
Expand Down Expand Up @@ -82,6 +85,7 @@
)
from dreadnode.user_config import UserConfig
from dreadnode.util import (
bump_version,
clean_str,
create_key_from_name,
handle_internal_errors,
Expand Down Expand Up @@ -1339,15 +1343,35 @@ def load_dataset(
def save_dataset_to_disk(
self,
ds: dataset.Dataset,
version: str | None = None,
strategy: VersionStrategy | None = None,
) -> None:
"""
Save a dataset to the local cache.

If a `version` or `strategy` is provided, the dataset's version will be
updated before saving.

Example:
```
dreadnode.save_dataset_to_disk(my_dataset)
```

Args:
ds: The dataset to save.
version: A specific version string to set for the dataset.
strategy: A versioning strategy to automatically bump the version
(e.g., 'major', 'minor', 'patch').
"""
if version is not None:
try:
parse_version(version)
except InvalidVersion as e:
raise ValueError(f"Invalid version string: {version}") from e
ds.metadata.version = version
elif strategy is not None:
new_version = bump_version(ds.metadata.version, strategy)
ds.metadata.version = new_version

dataset.save_dataset_to_disk(
dataset=ds,
Expand All @@ -1357,15 +1381,36 @@ def save_dataset_to_disk(
def push_dataset(
self,
ds: dataset.Dataset,
version: str | None = None,
strategy: VersionStrategy | None = None,
) -> None:
"""
Push a dataset to the Dreadnode server.

This will first save the dataset to the local cache and then upload it.
If a `version` or `strategy` is provided, the dataset's version will be
updated before pushing.

Example:
```
dreadnode.push_dataset(my_dataset)
```

Args:
ds: The dataset to push.
version: A specific version string to set for the dataset.
strategy: A versioning strategy to automatically bump the version
(e.g., 'major', 'minor', 'patch').
"""
if version is not None:
try:
parse_version(version)
except InvalidVersion as e:
raise ValueError(f"Invalid version string: {version}") from e
ds.metadata.version = version
elif strategy is not None:
new_version = bump_version(ds.metadata.version, strategy)
ds.metadata.version = new_version

dataset.push_dataset(
dataset=ds,
Expand Down
24 changes: 18 additions & 6 deletions dreadnode/storage/datasets/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,10 +145,23 @@ def get_remote_save_uri(self, metadata: DatasetMetadata) -> tuple[UUID, str]:
upload_request = CreateDatasetRequest(
org_key=metadata.organization,
key=metadata.name,
version=metadata.version,
tags=metadata.tags,
)

response = self._api.create_dataset(request=upload_request)
return response.dataset_id, response.user_data_access_response.uri
dataset_id = response.dataset_id
user_data_access_response = response.user_data_access_response
self._cached_s3_fs = pafs.S3FileSystem(
access_key=user_data_access_response.access_key_id,
secret_key=user_data_access_response.secret_access_key,
session_token=user_data_access_response.session_token,
endpoint_override=resolve_endpoint(user_data_access_response.endpoint),
region=user_data_access_response.region,
check_directory_existence_before_creation=True,
)
self._credentials_expiry = user_data_access_response.expiration
return dataset_id, user_data_access_response.uri

def remote_save_complete(self, dataset_id: str, *, complete: bool) -> None:
"""
Expand Down Expand Up @@ -184,9 +197,7 @@ def get_s3_config(self, dataset_id: UUID) -> dict[str, Any]:
if not self._api:
raise ValueError("No client configured")

creds = self._api.get_user_data_credentials(
organization_id=self.organization_id, dataset_id=dataset_id
)
creds = self._api.get_dataset_access_credentials(dataset_id=dataset_id)
self._credentials_expiry = creds.expiration
resolved_endpoint = resolve_endpoint(creds.endpoint)

Expand Down Expand Up @@ -224,13 +235,14 @@ def get_fs_and_path(self, uri: str) -> tuple[FileSystem, str]:

try:
_, path_body = uri.split("://", 1)
path_body = path_body.rstrip("/")
except ValueError:
return pafs.LocalFileSystem(), uri

if self._cached_s3_fs is None or self.needs_refresh():
try:
# Try to extract dataset ID from URI which expect is of the form dn://<org_id>/datasets/<dataset_id>
dataset_id = UUID(path_body.split("/")[-1])
# Try to extract dataset ID from URI which expect is of the form dn://<user bucket>/<org_id>/datasets/<dataset_id>/<version>
dataset_id = UUID(path_body.split("/")[3])
config = self.get_s3_config(dataset_id=dataset_id)
self._cached_s3_fs = pafs.S3FileSystem(**config)
except ValueError:
Expand Down
3 changes: 2 additions & 1 deletion dreadnode/storage/datasets/manifest.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from pathlib import Path
from typing import Any

import pyarrow as pa
from pyarrow.fs import FileSelector, FileSystem, FileType
from pydantic import BaseModel, Field
from tqdm import tqdm
Expand Down Expand Up @@ -134,7 +135,7 @@ def compute_file_hash(
with fs.open_input_stream(file_path) as f:
digest = hashlib.file_digest(f, algorithm)
return digest.hexdigest()
except Exception as e:
except (OSError, pa.ArrowException) as e:
logging_console.print(f"Failed to hash {file_path}: {e}")
return ""

Expand Down
56 changes: 56 additions & 0 deletions dreadnode/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,9 @@
warn_at_user_stacklevel as _warn_at_user_stacklevel,
)
from loguru import logger
from packaging.version import InvalidVersion, Version, parse

from dreadnode.common_types import VersionStrategy

P = ParamSpec("P")
R = TypeVar("R")
Expand Down Expand Up @@ -176,6 +179,17 @@ def valid_key(key: str) -> bool:
return bool(re.fullmatch(r"[a-z0-9-]+", key))


def valid_version(version: str) -> bool:
"""
Check if the version is valid (semantic versioning format).
"""
try:
parsed = parse(version)
return isinstance(parsed, Version)
except InvalidVersion:
return False


# Imports


Expand Down Expand Up @@ -901,3 +915,45 @@ def _wrapper(*args: P.args, **kwargs: P.kwargs) -> R:
def is_path_remote(uri: str) -> bool:
scheme = fsspec.utils.get_protocol(uri)
return scheme not in ("file", None, "")


# Versioning


def bump_version(version_str: str, strategy: VersionStrategy) -> str:
"""
Increments a version string based on the strategy.

Args:
version_str: The version string (e.g., "1.0.1", "2.3")
strategy: One of "major", "minor", "patch"

Returns:
The incremented version string.
"""
parsed = parse(version_str)

# Ensure it's a valid version object that has a release segment
if not isinstance(parsed, Version):
raise TypeError(f"Invalid version: {version_str}")

# properties: parsed.major, parsed.minor, parsed.micro
# Note: parsed.micro returns 0 even if not explicitly in the string (e.g. "1.2" -> micro=0)
major = parsed.major
minor = parsed.minor
micro = parsed.micro

if strategy == "major":
major += 1
minor = 0
micro = 0
elif strategy == "minor":
minor += 1
micro = 0
elif strategy == "patch":
micro += 1
else:
raise ValueError("Strategy must be 'major', 'minor', or 'patch'")

# Reconstruct the version string
return f"{major}.{minor}.{micro}"
Loading