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
160 changes: 140 additions & 20 deletions cli_tools/mcdi/mcdi/download/sources/pdc.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,37 @@
from __future__ import annotations

import base64
import binascii
import csv
import json
import time
import urllib.parse
from pathlib import Path
from typing import Optional

import requests

from ...errors import ApiError, InputError
from ...net import build_session
from .base import CANDIDATE_DELIMITERS, FileEntry, RateLimit, Source, read_header

# Columns present in every PDC file manifest (CSV/TSV) exported from the portal.
REQUIRED_HEADERS = {"PDC Study ID", "Data Category", "File Type", "File Download Link"}

GRAPHQL_URL = "https://pdc.cancer.gov/graphql"
_PAGE_SIZE = 100
_STALE_MARGIN_SECONDS = 900

FILES_PER_STUDY_QUERY = """
query FilesPerStudy($studyId: String!, $offset: Int!, $limit: Int!) {
filesPerStudy(pdc_study_id: $studyId, offset: $offset, limit: $limit) {
file_name
file_size
md5sum
signedUrl { url }
}
}
"""


def _find_md5_key(fieldnames: list[str]) -> str | None:
for name in fieldnames:
Expand All @@ -16,9 +40,66 @@ def _find_md5_key(fieldnames: list[str]) -> str | None:
return None


def _cloudfront_policy_expiry(url: str) -> Optional[int]:
query = urllib.parse.parse_qs(urllib.parse.urlparse(url).query)
policy = query.get("Policy", [None])[0]
if not policy:
return None
padded = policy.replace("-", "+").replace("_", "=").replace("~", "/")
padded += "=" * (-len(padded) % 4)
try:
data = json.loads(base64.b64decode(padded))
return int(data["Statement"][0]["Condition"]["DateLessThan"]["AWS:EpochTime"])
except (KeyError, IndexError, TypeError, ValueError, binascii.Error):
return None


def _is_stale(url: str) -> bool:
expiry = _cloudfront_policy_expiry(url)
return expiry is None or expiry <= time.time() + _STALE_MARGIN_SECONDS


class PdcUrlRefresher:
def __init__(self, session: Optional[requests.Session] = None):
self._session = session or build_session()

def _post(self, study_id: str, offset: int) -> list[dict]:
try:
resp = self._session.post(
GRAPHQL_URL,
json={
"query": FILES_PER_STUDY_QUERY,
"variables": {"studyId": study_id, "offset": offset, "limit": _PAGE_SIZE},
},
timeout=60,
)
except requests.RequestException as exc:
raise ApiError(f"PDC filesPerStudy request failed: {exc}") from exc
if resp.status_code >= 400:
raise ApiError(f"PDC filesPerStudy HTTP {resp.status_code}: {resp.text[:300]}")
payload = resp.json()
if "errors" in payload:
raise ApiError(f"PDC filesPerStudy({study_id!r}) errors: {payload['errors']}")
return payload["data"]["filesPerStudy"]

def fetch_study(self, study_id: str) -> dict[str, dict]:
by_name: dict[str, dict] = {}
offset = 0
while True:
page = self._post(study_id, offset)
for rec in page:
by_name[rec["file_name"]] = rec
if len(page) < _PAGE_SIZE:
return by_name
offset += _PAGE_SIZE


class PDCSource(Source):
name = "pdc"

def __init__(self, refresher: Optional[PdcUrlRefresher] = None):
self._refresher = refresher or PdcUrlRefresher()

@staticmethod
def sniff(header_fields: list[str]) -> bool:
return REQUIRED_HEADERS.issubset(set(header_fields))
Expand All @@ -28,34 +109,73 @@ def parse_manifest(self, path: Path) -> list[FileEntry]:
(d for d in CANDIDATE_DELIMITERS if self.sniff(read_header(path, d))),
CANDIDATE_DELIMITERS[-1],
)
entries = []
rows = []
with open(path, newline="") as f:
reader = csv.DictReader(f, delimiter=delimiter)
md5_key = _find_md5_key(reader.fieldnames or [])
for row in reader:
filename = row["File Name"].strip()
study_id = row["PDC Study ID"].strip()
study_version = row["PDC Study Version"].strip()
data_category = row["Data Category"].strip()
file_type = row["File Type"].strip()
run_metadata_id = (row.get("Run Metadata ID") or "").strip()
rows.append(
{
"filename": row["File Name"].strip(),
"study_id": row["PDC Study ID"].strip(),
"study_version": row["PDC Study Version"].strip(),
"data_category": row["Data Category"].strip(),
"file_type": row["File Type"].strip(),
"run_metadata_id": run_metadata_id,
"url": row["File Download Link"].strip(),
"md5": row[md5_key].strip() if md5_key and row.get(md5_key) else None,
}
)

parts = [study_id, study_version, data_category]
if run_metadata_id and run_metadata_id.lower() != "null":
parts.append(run_metadata_id)
parts.append(file_type)

entries.append(
FileEntry(
file_id=filename,
filename=filename,
rel_dir=Path("pdc").joinpath(*parts),
url=row["File Download Link"].strip(),
md5=row[md5_key].strip() if md5_key and row.get(md5_key) else None,
)
self._refresh_stale(rows)

entries = []
for row in rows:
parts = [row["study_id"], row["study_version"], row["data_category"]]
if row["run_metadata_id"] and row["run_metadata_id"].lower() != "null":
parts.append(row["run_metadata_id"])
parts.append(row["file_type"])
entries.append(
FileEntry(
file_id=row["filename"],
filename=row["filename"],
rel_dir=Path("pdc").joinpath(*parts),
url=row["url"],
md5=row["md5"],
)
)
return entries

def _refresh_stale(self, rows: list[dict]) -> None:
by_study: dict[str, list[dict]] = {}
for row in rows:
if _is_stale(row["url"]):
by_study.setdefault(row["study_id"], []).append(row)
if not by_study:
return

errors: list[str] = []
for study_id, stale_rows in by_study.items():
fresh = self._refresher.fetch_study(study_id)
for row in stale_rows:
rec = fresh.get(row["filename"])
if rec is None:
errors.append(f"{row['filename']!r} (study {study_id}): not found in current PDC data")
continue
fresh_md5 = rec.get("md5sum")
if row["md5"] and fresh_md5 and row["md5"].lower() != fresh_md5.lower():
errors.append(
f"{row['filename']!r} (study {study_id}): md5 mismatch "
f"(manifest {row['md5']!r} vs PDC {fresh_md5!r})"
)
continue
row["url"] = rec["signedUrl"]["url"]
row["md5"] = row["md5"] or fresh_md5

if errors:
raise InputError("PDC signed-URL refresh failed:\n" + "\n".join(errors))

def request_kwargs(self, entry: FileEntry) -> dict:
return {}

Expand Down
14 changes: 4 additions & 10 deletions cli_tools/mcdi/tests/test_download_pdc_network.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,14 +17,10 @@
STUDY_ID = "PDC000109"
SAMPLE_FILE_COUNT = 2

# `filesPerStudy` is the same query the PDC portal's Explore page uses to
# resolve each file's CloudFront `signedUrl`. It isn't in PDC's published
# API docs (found via GraphQL suggestion errors: `filesPerStudy(pdc_study_id)
# { signedUrl { url } }`), but it's a live query against the same public,
# tokenless GraphQL endpoint the portal itself calls, not a private one.
# https://pdc-docs.cancer.gov/pdc/publicapi-documentation#!/Files/filesPerStudy
FILES_PER_STUDY_QUERY = """
query FilesPerStudy($studyId: String!) {
filesPerStudy(pdc_study_id: $studyId) {
filesPerStudy(pdc_study_id: $studyId, offset: 0, limit: 100) {
file_id
file_name
file_size
Expand Down Expand Up @@ -53,11 +49,9 @@ def _fetch_sample_files() -> list[dict]:
)
resp.raise_for_status()
payload = resp.json()
if "errors" in payload:
pytest.skip(f"PDC GraphQL API returned errors: {payload['errors']}")
assert "errors" not in payload, payload.get("errors")
files = payload["data"]["filesPerStudy"]
if not files:
pytest.skip(f"PDC study {STUDY_ID} returned no files")
assert files, f"PDC study {STUDY_ID} returned no files"
return sorted(files, key=lambda f: int(f["file_size"]))[:SAMPLE_FILE_COUNT]


Expand Down
168 changes: 168 additions & 0 deletions cli_tools/mcdi/tests/test_download_pdc_refresh.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,168 @@
import base64
import json
import time

import pytest

from mcdi.download.sources.pdc import PDCSource, _cloudfront_policy_expiry, _is_stale
from mcdi.errors import InputError

STUDY = "PDC000001"


def _signed_url(expiry_epoch: int, resource: str = "https://x.cloudfront.net/f.raw") -> str:
policy = json.dumps(
{"Statement": [{"Resource": resource, "Condition": {"DateLessThan": {"AWS:EpochTime": expiry_epoch}}}]},
separators=(",", ":"),
).encode()
encoded = base64.b64encode(policy).decode().replace("+", "-").replace("=", "_").replace("/", "~")
return f"{resource}?Policy={encoded}&Key-Pair-Id=X&Signature=Y"


def _manifest(tmp_path, url, filename="f.raw", md5="abc123", study=STUDY):
path = tmp_path / "m.csv"
path.write_text(
"PDC Study ID,PDC Study Version,Data Category,File Type,File Name,File MD5sum,File Download Link\n"
f"{study},1,Raw Mass Spectra,raw,{filename},{md5},{url}\n"
)
return path


class FakeRefresher:
def __init__(self, studies: dict):
self._studies = studies
self.calls = []

def fetch_study(self, study_id):
self.calls.append(study_id)
return self._studies[study_id]


def test_cloudfront_policy_expiry_roundtrip():
url = _signed_url(1800000000)
assert _cloudfront_policy_expiry(url) == 1800000000


def test_cloudfront_policy_expiry_none_for_malformed_url():
assert _cloudfront_policy_expiry("https://example.com/f") is None
assert _cloudfront_policy_expiry("https://example.com/f?Policy=not-base64!!") is None


def test_is_stale_true_for_expired_and_unparseable():
assert _is_stale(_signed_url(1))
assert _is_stale("https://example.com/f")


def test_is_stale_false_for_far_future_expiry():
assert not _is_stale(_signed_url(int(time.time()) + 86400))


def test_fresh_url_skips_refresh(tmp_path):
url = _signed_url(int(time.time()) + 86400)
manifest = _manifest(tmp_path, url)

class Unused:
def fetch_study(self, study_id):
raise AssertionError("should not be called")

entries = PDCSource(refresher=Unused()).parse_manifest(manifest)
assert entries[0].url == url


def test_stale_url_refreshed(tmp_path):
stale_url = _signed_url(1)
fresh_url = _signed_url(int(time.time()) + 86400)
manifest = _manifest(tmp_path, stale_url, filename="f.raw", md5="abc123")
refresher = FakeRefresher({STUDY: {"f.raw": {"md5sum": "abc123", "signedUrl": {"url": fresh_url}}}})

entries = PDCSource(refresher=refresher).parse_manifest(manifest)

assert entries[0].url == fresh_url
assert refresher.calls == [STUDY]


def test_stale_url_missing_from_refresh_raises(tmp_path):
manifest = _manifest(tmp_path, _signed_url(1), filename="f.raw")
refresher = FakeRefresher({STUDY: {}})

with pytest.raises(InputError, match="f.raw"):
PDCSource(refresher=refresher).parse_manifest(manifest)


def test_stale_url_md5_mismatch_raises(tmp_path):
manifest = _manifest(tmp_path, _signed_url(1), filename="f.raw", md5="abc123")
refresher = FakeRefresher(
{STUDY: {"f.raw": {"md5sum": "different", "signedUrl": {"url": _signed_url(int(time.time()) + 86400)}}}}
)

with pytest.raises(InputError, match="md5 mismatch"):
PDCSource(refresher=refresher).parse_manifest(manifest)


def test_stale_url_no_manifest_md5_trusts_refresh(tmp_path):
stale_url = _signed_url(1)
fresh_url = _signed_url(int(time.time()) + 86400)
manifest = _manifest(tmp_path, stale_url, filename="f.raw", md5="")
refresher = FakeRefresher({STUDY: {"f.raw": {"md5sum": "fresh_md5", "signedUrl": {"url": fresh_url}}}})

entries = PDCSource(refresher=refresher).parse_manifest(manifest)

assert entries[0].url == fresh_url
assert entries[0].md5 == "fresh_md5"


def test_multiple_studies_grouped_one_call_each(tmp_path):
path = tmp_path / "m.csv"
stale = _signed_url(1)
fresh_a = _signed_url(int(time.time()) + 86400)
fresh_b = _signed_url(int(time.time()) + 86400)
path.write_text(
"PDC Study ID,PDC Study Version,Data Category,File Type,File Name,File MD5sum,File Download Link\n"
f"PDC000001,1,Raw Mass Spectra,raw,a.raw,md5a,{stale}\n"
f"PDC000002,1,Raw Mass Spectra,raw,b.raw,md5b,{stale}\n"
)
refresher = FakeRefresher(
{
"PDC000001": {"a.raw": {"md5sum": "md5a", "signedUrl": {"url": fresh_a}}},
"PDC000002": {"b.raw": {"md5sum": "md5b", "signedUrl": {"url": fresh_b}}},
}
)

entries = PDCSource(refresher=refresher).parse_manifest(path)

urls = {e.filename: e.url for e in entries}
assert urls == {"a.raw": fresh_a, "b.raw": fresh_b}
assert sorted(refresher.calls) == ["PDC000001", "PDC000002"]


def test_pagination_loops_until_short_page():
from mcdi.download.sources.pdc import _PAGE_SIZE, PdcUrlRefresher

class FakeSession:
def __init__(self, pages):
self.pages = pages
self.calls = 0

def post(self, url, json, timeout):
offset = json["variables"]["offset"]
page = self.pages[offset // _PAGE_SIZE]
self.calls += 1
return FakeResponse(page)

class FakeResponse:
def __init__(self, page):
self._page = page
self.status_code = 200

def json(self):
return {"data": {"filesPerStudy": self._page}}

full_page = [{"file_name": f"f{i}.raw", "md5sum": "x", "signedUrl": {"url": "u"}} for i in range(_PAGE_SIZE)]
last_page = [{"file_name": "last.raw", "md5sum": "x", "signedUrl": {"url": "u"}}]
session = FakeSession([full_page, last_page])

result = PdcUrlRefresher(session=session).fetch_study(STUDY)

assert session.calls == 2
assert len(result) == _PAGE_SIZE + 1
assert "last.raw" in result
Loading
Loading