diff --git a/docs/dd/tools/admet.md b/docs/dd/tools/admet.md index 2a69156d..9a0d07de 100644 --- a/docs/dd/tools/admet.md +++ b/docs/dd/tools/admet.md @@ -24,7 +24,7 @@ df = admet.run() from deeporigin.drug_discovery import LigandSet ligand_set = LigandSet.from_csv("ligands.csv") # "smiles" column -large = Admet(ligands=ligand_set) +large = Admet(ligands=ligand_set, batch_size=100) # optional; smaller batches run more in parallel (min 50) large.start() large.wait() df = large.get_results() diff --git a/src/drug_discovery/admet.py b/src/drug_discovery/admet.py index 9b86802b..c280d03c 100644 --- a/src/drug_discovery/admet.py +++ b/src/drug_discovery/admet.py @@ -22,12 +22,7 @@ from __future__ import annotations -import json -import os -from pathlib import Path -import tempfile from typing import Any, Literal, Self -import uuid from beartype import beartype import pandas as pd @@ -37,6 +32,11 @@ AsyncExecutableMixin, SyncExecutableMixin, ) +from deeporigin.drug_discovery.ligand_list_file import ( + ligand_rows_from_inputs, + ligands_from_rows, + upload_ligand_list, +) from deeporigin.drug_discovery.metabolism import _ligand_payloads from deeporigin.drug_discovery.notebook_watch_mixin import NotebookWatchMixin from deeporigin.drug_discovery.structures.ligand import Ligand, LigandSet @@ -47,10 +47,11 @@ from deeporigin.platform.project_scope import require_client_project_id from deeporigin.utils.constants import ( ADMET_EXECUTION_TIMEOUT_SECONDS, - ADMET_INLINE_LIGAND_CAP, ADMET_LIGAND_LIST_UPLOAD_PREFIX, + ADMET_MIN_BATCH_SIZE, ADMET_RESULT_EXPLORER_PAGE_SIZE, ADMET_WORKFLOW_LIGAND_THRESHOLD, + INLINE_LIGAND_CAP, QUOTE_APPROVE_AMOUNT, ) @@ -155,47 +156,6 @@ def _expand_admet_payload(payload: dict[str, Any]) -> list[dict[str, Any]]: return [] -def _ligands_from_payload_rows(raw: list[Any]) -> list[Ligand]: - """Rebuild ligands from inline or file JSON ligand rows.""" - - ligands: list[Ligand] = [] - for idx, row in enumerate(raw): - if not isinstance(row, dict): - raise ValueError( - f"Cannot rehydrate Admet: ligands[{idx}] is not an object." - ) - smiles = row.get("smiles") - if not smiles or not isinstance(smiles, str): - raise ValueError(f"Cannot rehydrate Admet: ligands[{idx}] has no SMILES.") - ligand = Ligand.from_smiles(smiles) - if row.get("id") is not None: - ligand.id = str(row["id"]) - ligands.append(ligand) - return ligands - - -def _ligands_from_list_file_bytes(payload: bytes) -> list[Ligand]: - """Parse a Ligand list file body into ligands.""" - - try: - text = payload.decode("utf-8") - except UnicodeDecodeError as exc: - raise ValueError( - f"Cannot rehydrate Admet: ligands_file is not valid UTF-8: {exc}" - ) from exc - try: - parsed = json.loads(text) - except json.JSONDecodeError as exc: - raise ValueError( - f"Cannot rehydrate Admet: ligands_file is not valid JSON: {exc.msg}" - ) from exc - if not isinstance(parsed, list) or not parsed: - raise ValueError( - "Cannot rehydrate Admet: ligands_file must be a non-empty JSON array." - ) - return _ligands_from_payload_rows(parsed) - - def _properties_from_inputs( inputs: dict[str, Any], ) -> tuple[str, ...] | None: @@ -233,31 +193,13 @@ def _ligands_from_inputs( if isinstance(project, dict) and project.get("id"): return [] - raw = inputs.get("ligands") - if isinstance(raw, list) and raw: - return _ligands_from_payload_rows(raw) - - remote = inputs.get("ligands_file") - if isinstance(remote, str) and remote.strip(): - if client is None or client.files is None: - raise ValueError( - "Cannot rehydrate Admet: client with files is required " - "to download ligands_file." - ) - try: - local_path = client.files.download(remote.strip(), direct=True) - payload = Path(local_path).read_bytes() - except Exception as exc: - raise ValueError( - f"Cannot rehydrate Admet: failed to download ligands_file " - f"{remote!r}: {exc}" - ) from exc - return _ligands_from_list_file_bytes(payload) - - raise ValueError( - "Cannot rehydrate Admet: stored inputs have no ligands, ligands_file, " - "or project." - ) + rows = ligand_rows_from_inputs(inputs, client=client, label="Admet") + if not rows: + raise ValueError( + "Cannot rehydrate Admet: stored inputs have no ligands, ligands_file, " + "or project." + ) + return ligands_from_rows(rows, label="Admet") class Admet( @@ -275,7 +217,7 @@ class Admet( ``run(quote=True)`` which sends SMILES without syncing first. Use :meth:`run` for at most - :data:`~deeporigin.utils.constants.ADMET_INLINE_LIGAND_CAP` ligands + :data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` ligands (blocking served path). For larger batches, call :meth:`start`, then :meth:`wait` or :meth:`watch`, then :meth:`get_results`. Project-wide runs use :class:`Admet` with ``ligands=[]`` and ``client.project_id`` set; call @@ -285,6 +227,8 @@ class Admet( ligands: Ligands whose SMILES are sent to the tool (empty for project runs). properties: Endpoint names for this run. method: Inference path — ``togo`` (default) or ``maplight``. + batch_size: Ligands per workflow pod on file and project runs, or + ``None`` for the tool default. """ tool_key: str = TOOL_KEYS_AND_VERSIONS["admet"]["tool_key"] @@ -296,6 +240,7 @@ def __init__( *, ligands: list[Ligand] | LigandSet, method: Literal["maplight", "togo"] = "togo", + batch_size: int | None = None, client: DeepOriginClient | None = None, ) -> None: """Configure an ADMET prediction run. @@ -303,7 +248,18 @@ def __init__( Fetches the live tool definition and fills :attr:`properties` with its endpoint enum. Pass ``ligands=[]`` for a project-wide workflow run (requires ``client.project_id``). + + ``batch_size`` sets ligands per workflow pod (``batchSize``) on file + and project runs; a smaller value fans out to more pods. It is ignored + on inline runs, which never fan out. ``None`` uses the tool default. + + Raises: + ValueError: If ``batch_size`` is below the tool minimum (50). """ + if batch_size is not None and batch_size < ADMET_MIN_BATCH_SIZE: + raise ValueError( + f"batch_size must be at least {ADMET_MIN_BATCH_SIZE} (got {batch_size})." + ) super().__init__(client=client) if isinstance(ligands, LigandSet): self._ligands: list[Ligand] = list(ligands.ligands) @@ -313,6 +269,7 @@ def __init__( self._allowed_endpoints: frozenset[str] | None = frozenset(endpoints) self._properties: list[str] | tuple[str, ...] | None = list(endpoints) self._method = method + self._batch_size = batch_size self._remote_ligands_file: str | None = None @property @@ -325,6 +282,11 @@ def method(self) -> str: """Selected admet-now inference method.""" return self._method + @property + def batch_size(self) -> int | None: + """Ligands per workflow pod, or ``None`` for the tool default.""" + return self._batch_size + def _is_project_run(self) -> bool: """True when this instance targets all ligands in ``client.project_id``.""" @@ -403,30 +365,15 @@ def _ensure_platform_inputs(self) -> None: LigandSet(ligands=self._ligands).sync(lazy=True, client=self.client) def _ensure_ligands_file_uploaded(self) -> str: - """Upload ligands JSON to UFA and return the remote path.""" + """Upload ligands JSON to UFA once and return the remote path.""" - if self._remote_ligands_file is not None: - return self._remote_ligands_file - if self.client.files is None: - raise ValueError( - "Cannot upload Ligand list file: client.files is not available." + if self._remote_ligands_file is None: + self._remote_ligands_file = upload_ligand_list( + _ligand_payloads(self._ligands), + client=self.client, + prefix=ADMET_LIGAND_LIST_UPLOAD_PREFIX, ) - - payloads = _ligand_payloads(self._ligands) - remote_path = f"{ADMET_LIGAND_LIST_UPLOAD_PREFIX}{uuid.uuid4().hex}.json" - fd, tmp_name = tempfile.mkstemp(suffix=".json", prefix="admet-ligands-") - try: - with os.fdopen(fd, "w", encoding="utf-8") as handle: - json.dump(payloads, handle, allow_nan=False) - self.client.files.upload(tmp_name, remote_path) - finally: - try: - os.unlink(tmp_name) - except OSError: - pass - - self._remote_ligands_file = remote_path - return remote_path + return self._remote_ligands_file def _make_inputs(self) -> dict[str, Any]: """Build tool ``inputs`` matching the admet-properties schema.""" @@ -436,6 +383,9 @@ def _make_inputs(self) -> dict[str, Any]: inputs["method"] = self._method if self._properties is not None: inputs["properties"] = list(self._properties) + workflow_run = self._is_project_run() or len(self._ligands) > INLINE_LIGAND_CAP + if workflow_run and self._batch_size is not None: + inputs["batchSize"] = self._batch_size if self._is_project_run(): project_id = require_client_project_id(self.client) @@ -443,7 +393,7 @@ def _make_inputs(self) -> dict[str, Any]: return inputs n = len(self._ligands) - if n > ADMET_INLINE_LIGAND_CAP: + if n > INLINE_LIGAND_CAP: remote = self._ensure_ligands_file_uploaded() inputs["ligands_file"] = remote inputs["ligands_count"] = n @@ -499,7 +449,7 @@ def _ensure_run_ligand_count(self) -> None: n = len(self._ligands) if n >= ADMET_WORKFLOW_LIGAND_THRESHOLD: raise ValueError( - f"run() supports at most {ADMET_INLINE_LIGAND_CAP} ligands " + f"run() supports at most {INLINE_LIGAND_CAP} ligands " f"(got {n}). Use start() then wait() or watch()." ) @@ -660,6 +610,8 @@ def from_dto( instance._allowed_endpoints = None method = inputs.get("method") instance._method = method if method in ("maplight", "togo") else "togo" + raw_batch = inputs.get("batchSize") + instance._batch_size = raw_batch if isinstance(raw_batch, int) else None remote = inputs.get("ligands_file") if isinstance(remote, str) and remote.strip(): instance._remote_ligands_file = remote.strip() diff --git a/src/drug_discovery/docking.py b/src/drug_discovery/docking.py index 0b4d1d45..ab0a2e50 100644 --- a/src/drug_discovery/docking.py +++ b/src/drug_discovery/docking.py @@ -19,6 +19,10 @@ AsyncExecutableMixin, SyncExecutableMixin, ) +from deeporigin.drug_discovery.ligand_list_file import ( + ligand_rows_from_inputs, + ligands_input, +) from deeporigin.drug_discovery.notebook_watch_mixin import NotebookWatchMixin from deeporigin.drug_discovery.structures.ligand import Ligand, LigandSet from deeporigin.drug_discovery.structures.pocket import Pocket @@ -398,8 +402,10 @@ def _build_tool_inputs( ) -> tuple[dict, dict]: """Build params and metadata for ``client.executions.create``. - Does not sync or upload; call :meth:`_ensure_platform_inputs` first when - inputs may not yet exist on the platform. + Does not sync; call :meth:`_ensure_platform_inputs` first when inputs may + not yet exist on the platform. Above + :data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` ligands, uploads a + Ligand list file. Args: ligand_set: Ligands to include in tool ``inputs`` (default: all @@ -427,7 +433,11 @@ def _build_tool_inputs( "id": self.protein.id, "file_path": self.protein.remote_path, }, - "ligands": [_ligand_tool_input_row(lig) for lig in ligands], + **ligands_input( + [_ligand_tool_input_row(lig) for lig in ligands], + client=self.client, + prefix="docking/ligand-lists/", + ), } return params, metadata @@ -513,10 +523,12 @@ def from_dto( "this execution may have been created with an older input schema." ) - ligands_input = inputs.get("ligands", []) - if not ligands_input: + ligand_rows = ligand_rows_from_inputs( + inputs, client=instance.client, label="Docking" + ) + if not ligand_rows: raise ValueError( - "Missing 'ligands' in execution userInputs; " + "Missing 'ligands' or 'ligands_file' in execution userInputs; " "this execution may have been created with an older input schema." ) @@ -530,10 +542,10 @@ def from_dto( ) fut_ligands = executor.submit( LigandSet.from_ids, - [lig["id"] for lig in ligands_input], + [lig["id"] for lig in ligand_rows], client=instance.client, download=False, - ligand_inputs=ligands_input, + ligand_inputs=ligand_rows, ) if pocket_id is not None: fut_pocket = executor.submit( diff --git a/src/drug_discovery/ligand_list_file.py b/src/drug_discovery/ligand_list_file.py new file mode 100644 index 00000000..a0c03597 --- /dev/null +++ b/src/drug_discovery/ligand_list_file.py @@ -0,0 +1,153 @@ +"""Ligand list files: inline ``ligands`` up to a cap, a UFA JSON upload above it. + +Admet, Metabolism, Docking and SecondaryPharmacology tools reject more than +:data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` inline ligands; larger +batches upload a Ligand list file (a bare JSON array of ligand rows) and pass +``ligands_file`` instead. +""" + +from __future__ import annotations + +import json +import os +from pathlib import Path +import tempfile +from typing import Any +import uuid + +from deeporigin.drug_discovery.structures.ligand import Ligand +from deeporigin.platform.client import DeepOriginClient +from deeporigin.utils.constants import INLINE_LIGAND_CAP + + +def upload_ligand_list( + rows: list[dict[str, Any]], + *, + client: DeepOriginClient, + prefix: str, +) -> str: + """Upload *rows* as a Ligand list file and return its UFA path. + + Args: + rows: Ligand rows in the tool's inline ``ligands`` shape. + client: Client whose ``files`` API does the upload. + prefix: UFA path prefix (e.g. ``"docking/ligand-lists/"``). + + Returns: + ``.json``. + """ + if client.files is None: + raise ValueError( + "Cannot upload Ligand list file: client.files is not available." + ) + remote_path = f"{prefix}{uuid.uuid4().hex}.json" + fd, tmp_name = tempfile.mkstemp(suffix=".json", prefix="ligand-list-") + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + json.dump(rows, handle, allow_nan=False) + client.files.upload(tmp_name, remote_path) + finally: + Path(tmp_name).unlink(missing_ok=True) + return remote_path + + +def ligands_input( + rows: list[dict[str, Any]], + *, + client: DeepOriginClient, + prefix: str, +) -> dict[str, Any]: + """Return ``{"ligands": rows}``, or upload *rows* and reference the file. + + Returns: + Tool inputs to merge: ``ligands``, or ``ligands_file`` + ``ligands_count``. + """ + if len(rows) <= INLINE_LIGAND_CAP: + return {"ligands": rows} + return { + "ligands_file": upload_ligand_list(rows, client=client, prefix=prefix), + "ligands_count": len(rows), + } + + +def parse_ligand_list(payload: bytes, *, label: str) -> list[dict[str, Any]]: + """Parse a Ligand list file body into rows. + + Raises: + ValueError: If the body is not UTF-8 JSON holding a non-empty array. + """ + try: + text = payload.decode("utf-8") + except UnicodeDecodeError as exc: + raise ValueError( + f"Cannot rehydrate {label}: ligands_file is not valid UTF-8: {exc}" + ) from exc + try: + parsed = json.loads(text) + except json.JSONDecodeError as exc: + raise ValueError( + f"Cannot rehydrate {label}: ligands_file is not valid JSON: {exc.msg}" + ) from exc + if not isinstance(parsed, list) or not parsed: + raise ValueError( + f"Cannot rehydrate {label}: ligands_file must be a non-empty JSON array." + ) + return parsed + + +def ligand_rows_from_inputs( + inputs: dict[str, Any], + *, + client: DeepOriginClient | None, + label: str, +) -> list[Any]: + """Return stored ligand rows from inline ``ligands`` or ``ligands_file``. + + Returns ``[]`` when neither is present; callers decide whether that is an + error. + + Raises: + ValueError: If ``ligands_file`` cannot be downloaded or parsed. + """ + raw = inputs.get("ligands") + if isinstance(raw, list) and raw: + return raw + remote = inputs.get("ligands_file") + if not isinstance(remote, str) or not remote.strip(): + return [] + if client is None or client.files is None: + raise ValueError( + f"Cannot rehydrate {label}: client with files is required " + "to download ligands_file." + ) + try: + local_path = client.files.download(remote.strip(), direct=True) + payload = Path(local_path).read_bytes() + except Exception as exc: + raise ValueError( + f"Cannot rehydrate {label}: failed to download ligands_file " + f"{remote!r}: {exc}" + ) from exc + return parse_ligand_list(payload, label=label) + + +def ligands_from_rows(raw: list[Any], *, label: str) -> list[Ligand]: + """Rebuild ligands from ``{smiles, id?}`` rows. + + Raises: + ValueError: If a row is not an object or has no SMILES. + """ + ligands: list[Ligand] = [] + for idx, row in enumerate(raw): + if not isinstance(row, dict): + raise ValueError( + f"Cannot rehydrate {label}: ligands[{idx}] is not an object." + ) + smiles = row.get("smiles") + if not smiles or not isinstance(smiles, str): + raise ValueError(f"Cannot rehydrate {label}: ligands[{idx}] has no SMILES.") + ligand = Ligand.from_smiles(smiles) + if row.get("id") is not None: + ligand.id = str(row["id"]) + ligands.append(ligand) + return ligands diff --git a/src/drug_discovery/metabolism.py b/src/drug_discovery/metabolism.py index 01e2ab42..a1acf4a9 100644 --- a/src/drug_discovery/metabolism.py +++ b/src/drug_discovery/metabolism.py @@ -20,7 +20,7 @@ :attr:`~deeporigin.drug_discovery.structures.ligand.Ligand.id` is already set. Batches larger than the platform Inline ligand cap -(:data:`~deeporigin.utils.constants.METABOLISM_INLINE_LIGAND_CAP`) dump a +(:data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP`) dump a Ligand list file to UFA and submit ``inputs.ligands_file`` on :meth:`start` (transparent to the caller). @@ -48,12 +48,7 @@ from __future__ import annotations -import json -import os -from pathlib import Path -import tempfile from typing import Any, Self -import uuid import warnings from beartype import beartype @@ -64,14 +59,19 @@ AsyncExecutableMixin, SyncExecutableMixin, ) +from deeporigin.drug_discovery.ligand_list_file import ( + ligand_rows_from_inputs, + ligands_from_rows, + upload_ligand_list, +) from deeporigin.drug_discovery.notebook_watch_mixin import NotebookWatchMixin from deeporigin.drug_discovery.structures.ligand import Ligand, LigandSet from deeporigin.exceptions import DeepOriginException from deeporigin.platform.client import DeepOriginClient from deeporigin.platform.constants import TOOL_KEYS_AND_VERSIONS, is_success_status from deeporigin.utils.constants import ( + INLINE_LIGAND_CAP, METABOLISM_EXECUTION_TIMEOUT_SECONDS, - METABOLISM_INLINE_LIGAND_CAP, METABOLISM_LIGAND_ID_QUERY_BATCH_SIZE, METABOLISM_RESULT_EXPLORER_PAGE_SIZE, METABOLISM_WORKFLOW_LIGAND_THRESHOLD, @@ -253,69 +253,6 @@ def _ligand_payloads(ligands: list[Ligand]) -> list[dict[str, str]]: return ligand_payloads -def _ligands_from_payload_rows(raw: list[Any]) -> list[Ligand]: - """Rebuild ligands from a bare Ligand list (inline or file JSON array). - - Args: - raw: List of ligand dicts with ``smiles`` and optional ``id``. - - Returns: - Ligands with SMILES and optional platform ids restored. - - Raises: - ValueError: If a row is not an object or has no SMILES. - """ - - ligands: list[Ligand] = [] - for idx, row in enumerate(raw): - if not isinstance(row, dict): - raise ValueError( - f"Cannot rehydrate Metabolism: ligands[{idx}] is not an object." - ) - smiles = row.get("smiles") - if not smiles or not isinstance(smiles, str): - raise ValueError( - f"Cannot rehydrate Metabolism: ligands[{idx}] has no SMILES." - ) - ligand = Ligand.from_smiles(smiles) - if row.get("id") is not None: - ligand.id = str(row["id"]) - ligands.append(ligand) - return ligands - - -def _ligands_from_list_file_bytes(payload: bytes) -> list[Ligand]: - """Parse a Ligand list file body into ligands. - - Args: - payload: UTF-8 JSON bytes whose root is a bare ligand array. - - Returns: - Parsed ligands. - - Raises: - ValueError: If the body is not valid UTF-8 JSON or not a ligand array. - """ - - try: - text = payload.decode("utf-8") - except UnicodeDecodeError as exc: - raise ValueError( - f"Cannot rehydrate Metabolism: ligands_file is not valid UTF-8: {exc}" - ) from exc - try: - parsed = json.loads(text) - except json.JSONDecodeError as exc: - raise ValueError( - f"Cannot rehydrate Metabolism: ligands_file is not valid JSON: {exc.msg}" - ) from exc - if not isinstance(parsed, list) or not parsed: - raise ValueError( - "Cannot rehydrate Metabolism: ligands_file must be a non-empty JSON array." - ) - return _ligands_from_payload_rows(parsed) - - def _ligands_from_inputs( inputs: dict[str, Any], *, @@ -338,27 +275,10 @@ def _ligands_from_inputs( row has no SMILES. """ - raw = inputs.get("ligands") - if isinstance(raw, list) and raw: - return _ligands_from_payload_rows(raw) - - remote = inputs.get("ligands_file") - if not isinstance(remote, str) or not remote.strip(): + rows = ligand_rows_from_inputs(inputs, client=client, label="Metabolism") + if not rows: raise ValueError("Cannot rehydrate Metabolism: stored inputs have no ligands.") - if client is None or client.files is None: - raise ValueError( - "Cannot rehydrate Metabolism: client with files is required " - "to download ligands_file." - ) - try: - local_path = client.files.download(remote.strip(), direct=True) - payload = Path(local_path).read_bytes() - except Exception as exc: - raise ValueError( - f"Cannot rehydrate Metabolism: failed to download ligands_file " - f"{remote!r}: {exc}" - ) from exc - return _ligands_from_list_file_bytes(payload) + return ligands_from_rows(rows, label="Metabolism") def _resolve_client(client: DeepOriginClient | None) -> DeepOriginClient: @@ -637,7 +557,7 @@ class Metabolism( :meth:`wait` or :meth:`watch`, then :meth:`get_results` / :meth:`get_molecules` (data platform first, ``jobOutputs`` fallback). Batches above - :data:`~deeporigin.utils.constants.METABOLISM_INLINE_LIGAND_CAP` upload a + :data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` upload a Ligand list file and pass ``ligands_file`` instead of inline ``ligands``. Use :meth:`fetch_results` / :meth:`fetch_molecules` to read indexed rows @@ -745,51 +665,30 @@ def fetch_molecules( ) def _ensure_ligands_file_uploaded(self) -> str: - """Dump ligands to JSON, upload to UFA, and return the remote path. + """Upload ligands JSON to UFA once and return the remote path. Caches the remote key on ``_remote_ligands_file`` so a repeated call does not re-upload. - - Returns: - UFA path under ``metabolism/ligand-lists/``. - - Raises: - ValueError: If the client has no files API. """ - if self._remote_ligands_file is not None: - return self._remote_ligands_file - if self.client.files is None: - raise ValueError( - "Cannot upload Ligand list file: client.files is not available." + if self._remote_ligands_file is None: + self._remote_ligands_file = upload_ligand_list( + _ligand_payloads(self._ligands), + client=self.client, + prefix=_LIGAND_LIST_UPLOAD_PREFIX, ) - - payloads = _ligand_payloads(self._ligands) - remote_path = f"{_LIGAND_LIST_UPLOAD_PREFIX}{uuid.uuid4().hex}.json" - fd, tmp_name = tempfile.mkstemp(suffix=".json", prefix="metabolism-ligands-") - try: - with os.fdopen(fd, "w", encoding="utf-8") as handle: - json.dump(payloads, handle, allow_nan=False) - self.client.files.upload(tmp_name, remote_path) - finally: - try: - os.unlink(tmp_name) - except OSError: - pass - - self._remote_ligands_file = remote_path - return remote_path + return self._remote_ligands_file def _make_inputs(self) -> dict[str, Any]: """Build tool ``inputs`` matching the metabolism schema. Batches larger than - :data:`~deeporigin.utils.constants.METABOLISM_INLINE_LIGAND_CAP` + :data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` upload a Ligand list file and return ``ligands_file``; smaller batches send inline ``ligands``. """ - if len(self._ligands) > METABOLISM_INLINE_LIGAND_CAP: + if len(self._ligands) > INLINE_LIGAND_CAP: return {"ligands_file": self._ensure_ligands_file_uploaded()} return {"ligands": _ligand_payloads(self._ligands)} diff --git a/src/drug_discovery/secondary_pharma.py b/src/drug_discovery/secondary_pharma.py index 4d233c05..38ce612e 100644 --- a/src/drug_discovery/secondary_pharma.py +++ b/src/drug_discovery/secondary_pharma.py @@ -50,12 +50,18 @@ AsyncExecutableMixin, SyncExecutableMixin, ) +from deeporigin.drug_discovery.ligand_list_file import ( + ligand_rows_from_inputs, + ligands_from_rows, + ligands_input, +) from deeporigin.drug_discovery.notebook_watch_mixin import NotebookWatchMixin from deeporigin.drug_discovery.structures.ligand import Ligand, LigandSet from deeporigin.drug_discovery.structures.pose import Pose, PoseSet from deeporigin.exceptions import DeepOriginException from deeporigin.platform.client import DeepOriginClient from deeporigin.platform.constants import TOOL_KEYS_AND_VERSIONS, is_success_status +from deeporigin.utils.constants import INLINE_LIGAND_CAP _UNIPROTS_ENUM_MISSING = ( "SecondaryPharmacology tool definition is missing a non-empty uniprots enum " @@ -162,31 +168,20 @@ def _ligand_ml_ligand_rows(ligands: list[Ligand]) -> list[dict[str, Any]]: return rows -def _ligands_from_inputs(inputs: dict[str, Any]) -> list[Ligand]: +def _ligands_from_inputs( + inputs: dict[str, Any], + *, + client: DeepOriginClient | None = None, +) -> list[Ligand]: """Rebuild ligands from stored secondary-pharma ``userInputs``. - Returns an empty list (rather than raising, unlike ``Admet``'s equivalent) - when ``ligands`` is absent -- a valid, expected shape for ``self_test`` runs. + Reads inline ``ligands``, or downloads ``ligands_file`` (raises + ``ValueError`` if it cannot be loaded). Returns an empty list (rather than + raising, unlike ``Admet``'s equivalent) when neither is present -- a valid, + expected shape for ``self_test`` runs. """ - raw = inputs.get("ligands") - if not isinstance(raw, list) or not raw: - return [] - ligands: list[Ligand] = [] - for idx, row in enumerate(raw): - if not isinstance(row, dict): - raise ValueError( - f"Cannot rehydrate SecondaryPharmacology: ligands[{idx}] is not an object." - ) - smiles = row.get("smiles") - if not smiles or not isinstance(smiles, str): - raise ValueError( - f"Cannot rehydrate SecondaryPharmacology: ligands[{idx}] has no SMILES." - ) - ligand = Ligand.from_smiles(smiles) - if row.get("id") is not None: - ligand.id = str(row["id"]) - ligands.append(ligand) - return ligands + rows = ligand_rows_from_inputs(inputs, client=client, label="SecondaryPharmacology") + return ligands_from_rows(rows, label="SecondaryPharmacology") # Result-explorer stores this table's rows under result_type="panelpose" -- @@ -904,11 +899,13 @@ def _ensure_method(self, expected: str, *, alternative_call: str) -> None: ) def _ensure_platform_inputs(self) -> None: - """Sync ligands to the data platform for the docking path. + """Sync ligands to the data platform so every row has a persisted id. - Only docking is a workflow that needs persisted ligand ids (as in - ``Docking._ensure_platform_inputs``). Ligand-ml scores straight from SMILES and - backfills ids afterwards; see :meth:`_backfill_ligand_ids`. + Docking always needs persisted ligand ids (as in + ``Docking._ensure_platform_inputs``), and so does any batch above + :data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` (a Ligand list + file). Small ligand-ml batches score straight from SMILES and backfill + ids afterwards; see :meth:`_backfill_ligand_ids`. """ LigandSet(ligands=self._ligands).sync(lazy=True, client=self.client) @@ -953,10 +950,16 @@ def _make_inputs(self) -> dict[str, Any]: "self_test": self._self_test, } if self._ligands: + if len(self._ligands) > INLINE_LIGAND_CAP: + # Preflight requires ``id`` on every Ligand list file row. + self._ensure_platform_inputs() if self._method == "docking": - inputs["ligands"] = [_docking_ligand_row(lig) for lig in self._ligands] + rows = [_docking_ligand_row(lig) for lig in self._ligands] else: - inputs["ligands"] = _ligand_ml_ligand_rows(self._ligands) + rows = _ligand_ml_ligand_rows(self._ligands) + inputs |= ligands_input( + rows, client=self.client, prefix="secondary-pharma/ligand-lists/" + ) if self._uniprots: inputs["uniprots"] = list(self._uniprots) return inputs @@ -1525,7 +1528,7 @@ def from_dto( inputs: dict[str, Any] = ( execution.get("userInputs") or execution.get("inputs") or {} ) - instance._ligands = _ligands_from_inputs(inputs) + instance._ligands = _ligands_from_inputs(inputs, client=instance.client) methods = inputs.get("methods") instance._method = ( methods[0] if isinstance(methods, list) and methods else "ligand-ml" diff --git a/src/utils/constants.py b/src/utils/constants.py index 8ec5cbca..87ac6e10 100644 --- a/src/utils/constants.py +++ b/src/utils/constants.py @@ -278,17 +278,17 @@ Cold-start model loading in the admet-now served image can exceed the default 600s POST timeout.""" -ADMET_INLINE_LIGAND_CAP = 100 -"""Max ligands sent inline in Admet ``inputs.ligands``. +INLINE_LIGAND_CAP = 100 +"""Max ligands sent inline in tool ``inputs.ligands``. -Matches the platform Inline ligand cap on ``deeporigin.admet-properties`` 2.x. -Larger batches upload a Ligand list file and pass ``ligands_file`` and -``ligands_count``.""" +Matches the platform Inline ligand cap on admet-properties 2.x, metabolism, +docking 4.x and secondary-pharma 3.x. Larger batches upload a Ligand list file +and pass ``ligands_file`` (see :mod:`deeporigin.drug_discovery.ligand_list_file`).""" ADMET_WORKFLOW_LIGAND_THRESHOLD = 101 """Ligand count at which ``Admet.run()`` must use ``start()`` instead. -``run()`` supports at most :data:`ADMET_INLINE_LIGAND_CAP` ligands (sync served). +``run()`` supports at most :data:`INLINE_LIGAND_CAP` ligands (sync served). At or above this threshold (101+), use ``start()`` then ``wait()`` / ``watch()``.""" ADMET_LIGAND_LIST_UPLOAD_PREFIX = "admet-properties/ligand-lists/" @@ -297,6 +297,9 @@ ADMET_RESULT_EXPLORER_PAGE_SIZE = 1000 """Page size when paging Admet ``admetproperty`` rows from result-explorer.""" +ADMET_MIN_BATCH_SIZE = 50 +"""Smallest ``batchSize`` (ligands per workflow pod) admet-properties accepts.""" + METABOLISM_WORKFLOW_LIGAND_THRESHOLD = 30 """Ligand count at which ``Metabolism.run()`` must use ``start()`` instead. @@ -304,12 +307,6 @@ workflows. ``run()`` raises for ``len(ligands) >=`` this value; use ``start()`` then ``wait()`` / ``watch()``.""" -METABOLISM_INLINE_LIGAND_CAP = 100 -"""Max ligands sent inline in Metabolism ``inputs.ligands``. - -Matches the platform Inline ligand cap. Larger batches dump a Ligand list -file to UFA and pass ``inputs.ligands_file`` instead.""" - METABOLISM_EXECUTION_TIMEOUT_SECONDS = 900.0 """HTTP timeout (seconds) for ``deeporigin.metabolism`` sync runs. diff --git a/tests/mock_server/routers/tools.py b/tests/mock_server/routers/tools.py index 0bad5aa9..c0bdf7d0 100644 --- a/tests/mock_server/routers/tools.py +++ b/tests/mock_server/routers/tools.py @@ -22,7 +22,7 @@ from fastapi import APIRouter, HTTPException, Request from deeporigin.utils.constants import ( - ADMET_INLINE_LIGAND_CAP, + INLINE_LIGAND_CAP, METABOLISM_WORKFLOW_LIGAND_THRESHOLD, ) @@ -3668,7 +3668,7 @@ async def run_tool( body.get("sync") is True and not has_file and not has_project - and n_ligands <= ADMET_INLINE_LIGAND_CAP + and n_ligands <= INLINE_LIGAND_CAP ) if sync_inline: execution = _build_admet_properties_execution( diff --git a/tests/test_admet_local.py b/tests/test_admet_local.py index 4aa698a2..26bea77d 100644 --- a/tests/test_admet_local.py +++ b/tests/test_admet_local.py @@ -10,8 +10,8 @@ from deeporigin.drug_discovery import Admet, Ligand from deeporigin.platform.constants import TOOL_KEYS_AND_VERSIONS from deeporigin.utils.constants import ( - ADMET_INLINE_LIGAND_CAP, ADMET_WORKFLOW_LIGAND_THRESHOLD, + INLINE_LIGAND_CAP, ) from tests.conftest import assert_quote_only_execution, check_tool_exists from tests.mock_server.routers.tools import ( @@ -320,7 +320,7 @@ def test_admet_start_above_inline_cap_uses_ligands_file( """Batches above the inline cap submit ``ligands_file`` and ``ligands_count``.""" _assert_tool_available(client) - n = ADMET_INLINE_LIGAND_CAP + 1 + n = INLINE_LIGAND_CAP + 1 ligands = [Ligand.from_smiles("CCO")] * n job = Admet(ligands=ligands, client=client) inputs = job._make_inputs() @@ -344,6 +344,32 @@ def test_admet_start_above_inline_cap_uses_ligands_file( assert prop in df.columns +def test_admet_batch_size_sent_on_workflow_paths_only( + client: DeepOriginClient, +) -> None: + """``batch_size`` becomes ``batchSize`` on file/project runs, not inline ones.""" + + _assert_tool_available(client) + many = [Ligand.from_smiles("CCO")] * (INLINE_LIGAND_CAP + 1) + assert Admet(ligands=many, client=client)._make_inputs().get("batchSize") is None + + file_job = Admet(ligands=many, batch_size=50, client=client) + assert file_job.batch_size == 50 + assert file_job._make_inputs()["batchSize"] == 50 + project_inputs = Admet(ligands=[], batch_size=60, client=client)._make_inputs() + assert project_inputs["batchSize"] == 60 + inline = Admet(ligands=[Ligand.from_smiles("CCO")], batch_size=50, client=client) + assert "batchSize" not in inline._make_inputs() + + with pytest.raises(ValueError, match="batch_size must be at least 50"): + Admet(ligands=many, batch_size=49, client=client) + + file_job.properties = list(_ADMET_PROPERTIES) + file_job.start() + restored = Admet.from_dto(file_job._dto or {}, client=client) + assert restored.batch_size == 50 + + def test_admet_run_rejects_workflow_scale_batch(client: DeepOriginClient) -> None: """``run()`` refuses 101+ ligands.""" diff --git a/tests/test_admet_unit.py b/tests/test_admet_unit.py index 42521186..2be66e67 100644 --- a/tests/test_admet_unit.py +++ b/tests/test_admet_unit.py @@ -8,7 +8,6 @@ _endpoints_from_definition, _execution_predictions, _ligands_from_inputs, - _ligands_from_list_file_bytes, _properties_from_inputs, _rows_from_result_explorer, _validate_admet_properties, @@ -148,29 +147,6 @@ def test_ligands_from_inputs_rejects_missing_rows() -> None: _ligands_from_inputs({"ligands": [{"id": "1"}]}) -def test_ligands_from_list_file_bytes_parses_rows() -> None: - """A Ligand list file rehydrates SMILES and platform ids.""" - ligands = _ligands_from_list_file_bytes(b'[{"smiles": "CCO", "id": 7}]') - assert [(lig.smiles, lig.id) for lig in ligands] == [("CCO", "7")] - - -@pytest.mark.parametrize( - ("payload", "match"), - [ - (b"\xff", "not valid UTF-8"), - (b"{", "not valid JSON"), - (b"[]", "non-empty JSON array"), - (b'{"smiles": "CCO"}', "non-empty JSON array"), - ], -) -def test_ligands_from_list_file_bytes_rejects_bad_bodies( - payload: bytes, match: str -) -> None: - """Corrupt list files raise ValueError instead of a parser exception.""" - with pytest.raises(ValueError, match=match): - _ligands_from_list_file_bytes(payload) - - def test_rows_from_result_explorer_flattens_and_skips_junk() -> None: """Nested, flat, and unrecognized result-explorer records.""" flat = {"ligand_id": "1", "hERG_classification": 0.2} diff --git a/tests/test_ligand_list_file.py b/tests/test_ligand_list_file.py new file mode 100644 index 00000000..826b2b51 --- /dev/null +++ b/tests/test_ligand_list_file.py @@ -0,0 +1,67 @@ +"""Unit tests for inline-vs-file ligand inputs.""" + +from pathlib import Path +import shutil +from types import SimpleNamespace + +import pytest + +from deeporigin.drug_discovery.ligand_list_file import ( + ligand_rows_from_inputs, + ligands_from_rows, + ligands_input, + parse_ligand_list, +) +from deeporigin.utils.constants import INLINE_LIGAND_CAP + + +class _Files: + def __init__(self, root: Path) -> None: + self.root = root + + def upload(self, local: str, remote: str) -> None: + dest = self.root / remote + dest.parent.mkdir(parents=True, exist_ok=True) + shutil.copy(local, dest) + + def download(self, remote: str, direct: bool = False) -> str: + return str(self.root / remote) + + +def test_ligands_inline_at_cap_and_file_above(tmp_path: Path) -> None: + """Exactly the cap stays inline; one more uploads and round-trips.""" + client = SimpleNamespace(files=_Files(tmp_path)) + rows = [{"id": f"L{i}", "smiles": "CCO"} for i in range(INLINE_LIGAND_CAP + 1)] + + assert ligands_input(rows[:-1], client=client, prefix="t/") == { + "ligands": rows[:-1] + } + + inputs = ligands_input(rows, client=client, prefix="t/") + assert set(inputs) == {"ligands_file", "ligands_count"} + assert inputs["ligands_count"] == INLINE_LIGAND_CAP + 1 + assert inputs["ligands_file"].startswith("t/") + assert ligand_rows_from_inputs(inputs, client=client, label="T") == rows + + +def test_ligands_from_rows_restores_smiles_and_ids() -> None: + """Rows rebuild ligands; numeric ids become strings, missing ids stay unset.""" + ligands = ligands_from_rows( + [{"smiles": "CCO", "id": 7}, {"smiles": "CCN"}], label="T" + ) + assert [(lig.smiles, lig.id) for lig in ligands] == [("CCO", "7"), ("CCN", None)] + + +@pytest.mark.parametrize( + ("payload", "match"), + [ + (b"\xff", "not valid UTF-8"), + (b"{", "not valid JSON"), + (b"[]", "non-empty JSON array"), + (b'{"smiles": "CCO"}', "non-empty JSON array"), + ], +) +def test_parse_ligand_list_rejects_bad_bodies(payload: bytes, match: str) -> None: + """Corrupt list files raise ValueError naming the caller.""" + with pytest.raises(ValueError, match=f"Cannot rehydrate T: .*{match}"): + parse_ligand_list(payload, label="T") diff --git a/tests/test_metabolism_local.py b/tests/test_metabolism_local.py index 9b12f6bf..bdbc8926 100644 --- a/tests/test_metabolism_local.py +++ b/tests/test_metabolism_local.py @@ -18,7 +18,7 @@ is_success_status, ) from deeporigin.utils.constants import ( - METABOLISM_INLINE_LIGAND_CAP, + INLINE_LIGAND_CAP, METABOLISM_WORKFLOW_LIGAND_THRESHOLD, ) from tests.conftest import check_tool_exists @@ -189,7 +189,7 @@ def test_metabolism_start_above_inline_cap_uses_ligands_file( ) -> None: """Batches above the inline cap submit ``ligands_file`` after UFA upload.""" _assert_tool_available(client) - n = METABOLISM_INLINE_LIGAND_CAP + 1 + n = INLINE_LIGAND_CAP + 1 ligands = [Ligand.from_smiles("CCO")] * n job = Metabolism(ligands=ligands, client=client) inputs = job._make_inputs() @@ -219,7 +219,7 @@ def test_metabolism_from_dto_rehydrates_ligands_file( ) -> None: """``from_dto`` downloads ``ligands_file`` and restores ligands.""" _assert_tool_available(client) - n = METABOLISM_INLINE_LIGAND_CAP + 1 + n = INLINE_LIGAND_CAP + 1 job = Metabolism(ligands=[Ligand.from_smiles("CCO")] * n, client=client) job.start() assert job.dto is not None diff --git a/tests/test_metabolism_unit.py b/tests/test_metabolism_unit.py index 60a29793..8b948592 100644 --- a/tests/test_metabolism_unit.py +++ b/tests/test_metabolism_unit.py @@ -19,7 +19,7 @@ ) from deeporigin.drug_discovery.structures.ligand import Ligand, LigandSet from deeporigin.utils.constants import ( - METABOLISM_INLINE_LIGAND_CAP, + INLINE_LIGAND_CAP, METABOLISM_RESULT_EXPLORER_PAGE_SIZE, METABOLISM_WORKFLOW_LIGAND_THRESHOLD, ) @@ -214,32 +214,9 @@ def test_ligands_from_inputs_requires_client_for_ligands_file() -> None: _ligands_from_inputs({"ligands_file": "metabolism/ligand-lists/x.json"}) -def test_ligands_from_list_file_bytes_parses_array() -> None: - """Bare JSON ligand arrays rehydrate into Ligand objects.""" - from deeporigin.drug_discovery.metabolism import _ligands_from_list_file_bytes - - raw = b'[{"smiles":"CCO","id":"lig-1"},{"smiles":"CCN"}]' - ligands = _ligands_from_list_file_bytes(raw) - assert [lig.smiles for lig in ligands] == ["CCO", "CCN"] - assert ligands[0].id == "lig-1" - assert ligands[1].id is None - - -def test_ligands_from_list_file_bytes_rejects_bad_json() -> None: - """Invalid UTF-8 or JSON fails with ValueError.""" - from deeporigin.drug_discovery.metabolism import _ligands_from_list_file_bytes - - with pytest.raises(ValueError, match="not valid UTF-8"): - _ligands_from_list_file_bytes(b"\xff\xfe") - with pytest.raises(ValueError, match="not valid JSON"): - _ligands_from_list_file_bytes(b"{not-json") - with pytest.raises(ValueError, match="non-empty JSON array"): - _ligands_from_list_file_bytes(b"{}") - - def test_make_inputs_stays_inline_at_cap() -> None: """Exactly the inline cap still sends ``ligands[]``.""" - n = METABOLISM_INLINE_LIGAND_CAP + n = INLINE_LIGAND_CAP job = Metabolism(ligands=[Ligand.from_smiles("CCO")] * n) inputs = job._make_inputs() assert "ligands_file" not in inputs diff --git a/tests/test_secondary_pharma.py b/tests/test_secondary_pharma.py index b20fba61..685e8c6f 100644 --- a/tests/test_secondary_pharma.py +++ b/tests/test_secondary_pharma.py @@ -31,6 +31,7 @@ Ligand, SecondaryPharmacology, ) +from deeporigin.drug_discovery.ligand_list_file import ligand_rows_from_inputs from deeporigin.drug_discovery.structures.pose import Pose, PoseSet from deeporigin.exceptions import DeepOriginException from deeporigin.platform.constants import ( @@ -39,6 +40,7 @@ is_success_status, ) from deeporigin.plots import WHITE_RED_HAZARD_PALETTE +from deeporigin.utils.constants import INLINE_LIGAND_CAP from tests.conftest import check_tool_exists from tests.mock_server.routers.tools import ( MOCK_SECONDARY_PHARMA_PANEL, @@ -442,6 +444,25 @@ def test_secondary_pharma_make_inputs_docking_uses_ligand_id_directly( assert inputs["ligands"] == [{"id": "manually-set-id", "smiles": "CCO"}] +def test_secondary_pharma_make_inputs_ligand_ml_above_cap_uses_synced_file( + client: DeepOriginClient, +) -> None: + """Above the inline cap, ligand-ml syncs ligands and sends a Ligand list file. + + Preflight requires ``id`` on every file row, so the unsynced-ligand-ml + shortcut above doesn't apply here. + """ + _assert_tool_available(client) + ligands = [Ligand.from_smiles("C" * (i + 1)) for i in range(INLINE_LIGAND_CAP + 1)] + job = SecondaryPharmacology(ligands=ligands, method="ligand-ml", client=client) + inputs = job._make_inputs() + assert "ligands" not in inputs + assert inputs["ligands_count"] == INLINE_LIGAND_CAP + 1 + rows = ligand_rows_from_inputs(inputs, client=client, label="T") + assert len(rows) == INLINE_LIGAND_CAP + 1 + assert all(row.get("id") for row in rows) + + def test_secondary_pharma_ensure_platform_inputs_syncs_ligands( client: DeepOriginClient, ) -> None: