diff --git a/databusclient/api/download.py b/databusclient/api/download.py index 356a2dd..cb24c02 100644 --- a/databusclient/api/download.py +++ b/databusclient/api/download.py @@ -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 @@ -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: @@ -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. @@ -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: @@ -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. @@ -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: @@ -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 @@ -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}" ) @@ -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 @@ -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) @@ -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 @@ -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] @@ -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 @@ -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 @@ -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 @@ -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": @@ -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, @@ -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: @@ -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: @@ -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: @@ -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: @@ -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", []) @@ -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. @@ -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 @@ -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, @@ -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 diff --git a/tests/test_download.py b/tests/test_download.py index b66521a..fcfd16a 100644 --- a/tests/test_download.py +++ b/tests/test_download.py @@ -1,11 +1,14 @@ """Download Tests""" import bz2 +import json import pytest from click.testing import CliRunner - -from databusclient.api.download import download as api_download +from databusclient.api.download import _parse_version_key +from databusclient.api.download import ( + download as api_download, +) from databusclient.cli import app # TODO: overall test structure not great, needs refactoring @@ -29,7 +32,10 @@ class FakeHeadResp: status_code = 200 - headers = {} + + # headers = {} + def __init__(self): + self.headers = {} class FakeGetResp: @@ -68,23 +74,100 @@ def test_with_collection(): api_download("tmp", DEFAULT_ENDPOINT, [TEST_COLLECTION]) +@pytest.mark.parametrize( + "versions, expected", + [ + ( + [ + "https://databus.dbpedia.org/example/2026-09-17", + "https://databus.dbpedia.org/example/2026-09-18", + "https://databus.dbpedia.org/example/2025-12-31", + ], + [ + "https://databus.dbpedia.org/example/2026-09-18", + "https://databus.dbpedia.org/example/2026-09-17", + "https://databus.dbpedia.org/example/2025-12-31", + ], + ), + ( + [ + "https://databus.dbpedia.org/example/2026.09.17", + "https://databus.dbpedia.org/example/2026.10.01", + "https://databus.dbpedia.org/example/2026.09.05", + ], + [ + "https://databus.dbpedia.org/example/2026.10.01", + "https://databus.dbpedia.org/example/2026.09.17", + "https://databus.dbpedia.org/example/2026.09.05", + ], + ), + ( + [ + "https://databus.dbpedia.org/example/2.9", + "https://databus.dbpedia.org/example/2.10", + "https://databus.dbpedia.org/example/2.1", + ], + [ + "https://databus.dbpedia.org/example/2.10", + "https://databus.dbpedia.org/example/2.9", + "https://databus.dbpedia.org/example/2.1", + ], + ), + ( + [ + "https://databus.dbpedia.org/example/2.9.0", + "https://databus.dbpedia.org/example/2.10.0", + "https://databus.dbpedia.org/example/2.1.0", + ], + [ + "https://databus.dbpedia.org/example/2.10.0", + "https://databus.dbpedia.org/example/2.9.0", + "https://databus.dbpedia.org/example/2.1.0", + ], + ), + ( + [ + "https://databus.dbpedia.org/example/v2.9.0", + "https://databus.dbpedia.org/example/v2.10.0", + "https://databus.dbpedia.org/example/v2.1.0", + ], + [ + "https://databus.dbpedia.org/example/v2.10.0", + "https://databus.dbpedia.org/example/v2.9.0", + "https://databus.dbpedia.org/example/v2.1.0", + ], + ), + ], +) +def test_parse_version_key_sorting(versions, expected): + assert sorted(versions, key=_parse_version_key, reverse=True) == expected + + def test_404_records_failed_manifest_entry(monkeypatch): - from databusclient.manifest.context import ManifestContext import databusclient.api.download as dl + from databusclient.manifest.context import ManifestContext class FakeGetResp: status_code = 404 - headers = {"content-length": "0"} + + # headers = {"content-length": "0"} + def __init__(self): + self.headers = {"content-length": "0"} def raise_for_status(self): import requests + raise requests.exceptions.HTTPError(response=self) monkeypatch.setattr("requests.head", lambda *a, **k: FakeHeadResp()) monkeypatch.setattr("requests.get", lambda *a, **k: FakeGetResp()) ctx = ManifestContext(command="download") - dl._download_file("https://databus.dbpedia.org/account/notexisting", localDir=".", manifest_context=ctx) + dl._download_file( + "https://databus.dbpedia.org/account/notexisting", + localDir=".", + manifest_context=ctx, + ) assert len(ctx.files) == 1 assert ctx.files[0]["status"] == "failed" @@ -137,7 +220,10 @@ def test_failed_download_does_not_create_graph_sidecar(monkeypatch, tmp_path): class FakeGet404Resp: status_code = 404 - headers = {"content-length": "0"} + + # headers = {"content-length": "0"} + def __init__(self): + self.headers = {"content-length": "0"} def raise_for_status(self): import requests @@ -168,8 +254,6 @@ def test_invalid_graph_mode_value_is_rejected_by_cli(): def test_cli_manifest_records_graph_mode(monkeypatch, tmp_path): - import json - captured = {} def fake_download(**kwargs):