Skip to content
Open
102 changes: 64 additions & 38 deletions databusclient/api/download.py
Original file line number Diff line number Diff line change
@@ -1,31 +1,30 @@
import json
import os
import bz2
import gzip
import json
import lzma
from typing import List, Optional, Tuple
import os
import re
import shutil
import tempfile
from datetime import UTC, datetime
from urllib.parse import urlparse

import requests
from SPARQLWrapper import JSON, SPARQLWrapper
from tqdm import tqdm
from datetime import datetime, timezone

from databusclient.api.utils import (
compute_sha256_and_length,
fetch_databus_jsonld,
get_databus_id_parts_from_file_url,
compute_sha256_and_length,
)
from databusclient.filehandling.format import (
FORMAT_TO_EXTENSION,
convert_file,
detect_format_from_filename,
get_converted_filename,
normalize_format,
get_format_class,
detect_format_from_filename,
FORMAT_TO_EXTENSION,
normalize_format,
)

# Compression format mappings
Expand Down Expand Up @@ -94,7 +93,7 @@ def _collect_files(directory: str) -> list[str]:
)


def _detect_compression_format(filename: str) -> Optional[str]:
def _detect_compression_format(filename: str) -> str | None:
"""Detect compression format from file extension.

Args:
Expand All @@ -111,8 +110,8 @@ def _detect_compression_format(filename: str) -> Optional[str]:


def _should_convert_compression(
filename: str, compression: Optional[str]
) -> Tuple[bool, Optional[str]]:
filename: str, compression: str | None
) -> tuple[bool, str | None]:
"""Determine if a file should have its compression format converted or compressed.

Source compression is detected automatically from the file extension.
Expand Down Expand Up @@ -213,7 +212,9 @@ def _convert_compression_format(

# Decompression-only path: target_format == 'none'
if target_format.lower() == "none":
print(f"Decompressing {os.path.basename(source_file)} -> {os.path.basename(target_file)}")
print(
f"Decompressing {os.path.basename(source_file)} -> {os.path.basename(target_file)}"
)
try:
with source_module.open(source_file, "rb") as sf:
with open(target_file, "wb") as tf:
Expand Down Expand Up @@ -371,7 +372,7 @@ def _extract_checksums_from_jsonld(json_str: str) -> dict:
return checksums


def _resolve_checksums_for_urls(file_urls: List[str], databus_key: str | None) -> dict:
def _resolve_checksums_for_urls(file_urls: list[str], databus_key: str | None) -> dict:
"""
Group file URLs by their Version URI, fetch each Version JSON-LD once,
and return a combined url->checksum mapping for the provided URLs.
Expand All @@ -381,7 +382,9 @@ def _resolve_checksums_for_urls(file_urls: List[str], databus_key: str | None) -
versions_map: dict = {}
for file_url in file_urls:
try:
host, accountId, groupId, artifactId, versionId, fileId = get_databus_id_parts_from_file_url(file_url)
host, accountId, groupId, artifactId, versionId, fileId = (
get_databus_id_parts_from_file_url(file_url)
)
except Exception:
continue
if versionId is None:
Expand Down Expand Up @@ -586,14 +589,14 @@ def _download_file(

# --- 5. Verify download size ---
if total_size_in_bytes != 0 and progress_bar.n != total_size_in_bytes:
raise IOError("Downloaded size does not match Content-Length header")
raise OSError("Downloaded size does not match Content-Length header")

# --- 6. Validate checksum on original downloaded file (BEFORE conversion) ---
actual_checksum = None
if validate_checksum:
try:
actual_checksum, _ = compute_sha256_and_length(filename)
except (OSError, IOError) as e:
except OSError as e:
print(f"WARNING: error computing checksum for {filename}: {e}")
actual_checksum = None

Expand All @@ -611,7 +614,7 @@ def _download_file(
os.remove(filename)
except OSError:
pass
raise IOError(
raise OSError(
f"Checksum mismatch for {filename}: expected {expected_checksum}, got {actual_checksum}"
)

Expand All @@ -630,7 +633,7 @@ def _download_file(
status="success",
sha256=actual_checksum or expected_checksum,
size_bytes=total_size_in_bytes if total_size_in_bytes else None,
downloaded_at=datetime.now(timezone.utc).isoformat(),
downloaded_at=datetime.now(UTC).isoformat(),
)
return

Expand All @@ -656,7 +659,9 @@ def _download_file(
# Decompress — strip compression extension, save plain file.
target_filename = _get_converted_filename(file, source_fmt, "none")
target_filepath = os.path.join(localDir, target_filename)
_convert_compression_format(filename, target_filepath, source_fmt, "none")
_convert_compression_format(
filename, target_filepath, source_fmt, "none"
)
else:
target_filename = _get_converted_filename(file, source_fmt, compression)
target_filepath = os.path.join(localDir, target_filename)
Expand All @@ -673,7 +678,7 @@ def _download_file(
status="success",
sha256=actual_checksum or expected_checksum,
size_bytes=total_size_in_bytes if total_size_in_bytes else None,
downloaded_at=datetime.now(timezone.utc).isoformat(),
downloaded_at=datetime.now(UTC).isoformat(),
)
return

Expand All @@ -690,7 +695,9 @@ def _download_file(
file, source_fmt, compression
)
target_filepath = os.path.join(localDir, target_filename)
_convert_compression_format(filename, target_filepath, source_fmt, compression)
_convert_compression_format(
filename, target_filepath, source_fmt, compression
)
final_paths = [target_filepath]
else:
final_paths = [filename]
Expand All @@ -702,7 +709,7 @@ def _download_file(
status="success",
sha256=actual_checksum or expected_checksum,
size_bytes=total_size_in_bytes if total_size_in_bytes else None,
downloaded_at=datetime.now(timezone.utc).isoformat(),
downloaded_at=datetime.now(UTC).isoformat(),
)
return

Expand Down Expand Up @@ -741,9 +748,12 @@ def _download_file(
source_format_for_mapping = detect_format_from_filename(conversion_input_path)
source_class_for_mapping = (
get_format_class(source_format_for_mapping)
if source_format_for_mapping else None
if source_format_for_mapping
else None
)
is_quad_to_triple = (
source_class_for_mapping == "quads" and target_class == "triples"
)
is_quad_to_triple = (source_class_for_mapping == "quads" and target_class == "triples")

if is_quad_to_triple:
# Output directory name = original filename with compression and
Expand Down Expand Up @@ -775,7 +785,7 @@ def _download_file(
status="success",
sha256=actual_checksum or expected_checksum,
size_bytes=total_size_in_bytes if total_size_in_bytes else None,
downloaded_at=datetime.now(timezone.utc).isoformat(),
downloaded_at=datetime.now(UTC).isoformat(),
)
return

Expand Down Expand Up @@ -807,7 +817,9 @@ def _download_file(
if source_compression is not None:
if should_convert_compression and compression:
# 'none' means no recompression after format conversion
final_compression = None if compression.lower() == "none" else compression
final_compression = (
None if compression.lower() == "none" else compression
)
else:
final_compression = source_compression
elif compression and compression.lower() != "none":
Expand Down Expand Up @@ -847,12 +859,12 @@ def _download_file(
status="success",
sha256=actual_checksum or expected_checksum,
size_bytes=total_size_in_bytes if total_size_in_bytes else None,
downloaded_at=datetime.now(timezone.utc).isoformat(),
downloaded_at=datetime.now(UTC).isoformat(),
)


def _download_files(
urls: List[str],
urls: list[str],
localDir: str,
vault_token_file: str = None,
databus_key: str = None,
Expand Down Expand Up @@ -948,7 +960,7 @@ def _query_sparql_endpoint(endpoint_url, query, databus_key=None) -> dict:

def _get_file_download_urls_from_sparql_query(
endpoint_url, query, databus_key=None
) -> List[str]:
) -> list[str]:
"""Execute a SPARQL query to get databus file download URLs.

Args:
Expand All @@ -965,7 +977,7 @@ def _get_file_download_urls_from_sparql_query(
if not isinstance(bindings, list):
raise ValueError("Invalid SPARQL response: 'bindings' missing or not a list")

urls: List[str] = []
urls: list[str] = []

for binding in bindings:
if not isinstance(binding, dict) or len(binding) != 1:
Expand Down Expand Up @@ -1224,9 +1236,21 @@ def _download_artifact(
)


def _parse_version_key(url: str) -> tuple:
"""Return a numeric sort key derived from the version segment of a Databus URL.

Splits the trailing version segment by non-digit characters and compares each
part as an integer, so '2.10.0' correctly sorts after '2.9.0' (unlike plain
lexicographic sort where '2.9' > '2.10' as strings).
"""
segment = url.rstrip("/").split("/")[-1]
parts = re.split(r"[^0-9]+", segment)
return tuple(int(part) for part in parts if part.isdigit())


def _get_databus_versions_of_artifact(
json_str: str, all_versions: bool
) -> str | List[str]:
) -> str | list[str]:
"""Parse the JSON-LD of a databus artifact to extract URLs of its versions.

Args:
Expand Down Expand Up @@ -1255,14 +1279,14 @@ def _get_databus_versions_of_artifact(
if not version_urls:
raise ValueError("No versions found in artifact JSON-LD")

version_urls.sort(reverse=True) # Sort versions in descending order
version_urls.sort(key=_parse_version_key, reverse=True)

if all_versions:
return version_urls
return version_urls[0]


def _get_file_download_urls_from_artifact_jsonld(json_str: str) -> List[str]:
def _get_file_download_urls_from_artifact_jsonld(json_str: str) -> list[str]:
"""Parse the JSON-LD of a databus artifact version to extract download URLs.

Args:
Expand All @@ -1272,7 +1296,7 @@ def _get_file_download_urls_from_artifact_jsonld(json_str: str) -> List[str]:
List of all file download URLs in the artifact version.
"""

databusIdUrl: List[str] = []
databusIdUrl: list[str] = []

json_dict = json.loads(json_str)
graph = json_dict.get("@graph", [])
Expand Down Expand Up @@ -1340,7 +1364,7 @@ def _download_group(
)


def _get_databus_artifacts_of_group(json_str: str) -> List[str]:
def _get_databus_artifacts_of_group(json_str: str) -> list[str]:
"""
Parse the JSON-LD of a databus group to extract URLs of all artifacts.

Expand All @@ -1361,7 +1385,7 @@ def _get_databus_artifacts_of_group(json_str: str) -> List[str]:
f"Unexpected type for 'databus:hasArtifact': {type(artifacts).__name__}"
)

result: List[str] = []
result: list[str] = []
for item in artifacts_iter:
if not isinstance(item, dict):
continue
Expand All @@ -1377,7 +1401,7 @@ def _get_databus_artifacts_of_group(json_str: str) -> List[str]:
def download(
localDir: str,
endpoint: str,
databusURIs: List[str],
databusURIs: list[str],
token=None,
databus_key=None,
all_versions=None,
Expand Down Expand Up @@ -1413,7 +1437,9 @@ def download(
"""
_validate_graph_mode(graph_mode)
for databusURI in databusURIs:
host, account, group, artifact, version, file = get_databus_id_parts_from_file_url(databusURI)
host, account, group, artifact, version, file = (
get_databus_id_parts_from_file_url(databusURI)
)

# Determine endpoint per-URI if not explicitly provided
uri_endpoint = endpoint
Expand Down
Loading