← Files Biohub ESMARCHIVED FILE

scripts/biohub_esm_lib/managed.py

57.4 KB · Sep 30, 2026 · 23:14 UTC

↓ Download file

"""Managed Biohub response normalization and structure artifacts."""

from __future__ import annotations

import math
from collections import Counter
from copy import deepcopy
from pathlib import Path
from typing import Any, Callable

from .constants import ESM_GIT_REVISION
from .errors import SchemaDriftError, ValidationError
from .managed_readiness import load_structure_serializer_sdk
from .provenance import (
    artifact_record,
    numeric_metric_summary,
    verify_installed_vcs_revision,
    write_bytes_atomic,
)

MANAGED_CONFIDENCE_FIELDS = (
    ("plddt", "plddt"),
    ("mean_plddt", "mean_plddt"),
    ("ptm", "ptm"),
    ("iptm", "iptm"),
    ("interface_ptm", "iptm"),
    ("pae", "pae"),
    ("pair_chains_iptm", "pair_chains_iptm"),
)

MANAGED_REQUESTED_OUTPUT_FIELDS = {
    "include_pae": ("pae",),
    "include_pair_chains_iptm": ("pair_chains_iptm",),
    "include_distogram": ("distogram",),
    "include_embeddings": (
        "output_embedding_sequence",
        "output_embedding_pair_pooled",
    ),
}

PLDDT_TRACK_REL_TOLERANCE = 1e-5
PLDDT_TRACK_ABS_TOLERANCE = 5e-4
COVALENT_BOND_SEVERE_MIN_ANGSTROM = 0.9
COVALENT_BOND_SEVERE_MAX_ANGSTROM = 2.5

PROTEIN_RESIDUE_TOKENS = {
    "A": "ALA",
    "R": "ARG",
    "N": "ASN",
    "D": "ASP",
    "C": "CYS",
    "Q": "GLN",
    "E": "GLU",
    "G": "GLY",
    "H": "HIS",
    "I": "ILE",
    "L": "LEU",
    "K": "LYS",
    "M": "MET",
    "F": "PHE",
    "P": "PRO",
    "S": "SER",
    "T": "THR",
    "W": "TRP",
    "Y": "TYR",
    "V": "VAL",
    "X": "UNK",
}
DNA_RESIDUE_TOKENS = {"A": "DA", "T": "DT", "C": "DC", "G": "DG"}
RNA_RESIDUE_TOKENS = {"A": "A", "U": "U", "C": "C", "G": "G"}


def _managed_confidence_containers(result: dict[str, Any]) -> list[dict[str, Any]]:
    """Return every named confidence container once, including nested envelopes."""

    containers: list[dict[str, Any]] = []
    pending = [result]
    seen: set[int] = set()
    while pending:
        container = pending.pop()
        identity = id(container)
        if identity in seen:
            continue
        seen.add(identity)
        containers.append(container)
        pending.extend(
            value
            for key in ("result", "confidence", "metrics")
            if isinstance((value := container.get(key)), dict)
        )
    return containers


def normalize_managed_response(response: dict[str, Any]) -> dict[str, Any]:
    """Apply the pinned SDK's single optional data envelope."""

    if "data" in response and "outputs" not in response:
        data = response["data"]
        if not isinstance(data, dict):
            raise SchemaDriftError(
                "managed API data envelope must contain a JSON object",
                raw=response,
            )
        return data
    return response


def _expected_confidence_dimensions(
    result: dict[str, Any],
    *,
    endpoint: str | None,
    request: dict[str, Any] | None,
) -> tuple[int | None, int | None, str | None, int | None]:
    if endpoint in {"fold", "fold_all_atom"} and request is not None:
        sequence = request.get("sequence")
        if isinstance(sequence, str) and sequence:
            return (
                sum(residue != "|" for residue in sequence),
                len(sequence),
                sequence,
                sequence.count("|") + 1,
            )
    request_chain_count: int | None = None
    if endpoint == "fold_all_atom" and request is not None:
        all_atom_input = request.get("all_atom_input")
        if isinstance(all_atom_input, dict):
            entities = all_atom_input.get("sequences")
            if isinstance(entities, list) and entities:
                request_chain_count = 0
                for entity in entities:
                    if not isinstance(entity, dict):
                        request_chain_count = None
                        break
                    chain_ids = entity.get("id")
                    request_chain_count += len(chain_ids) if isinstance(chain_ids, list) else 1
    coordinates = result.get("coordinates")
    if isinstance(coordinates, list) and coordinates:
        return len(coordinates), len(coordinates), None, None
    complex_state = result.get("complex")
    if isinstance(complex_state, dict):
        sequence = complex_state.get("sequence")
        if isinstance(sequence, list) and sequence:
            return len(sequence), len(sequence), None, request_chain_count
    return None, None, None, request_chain_count


def _require_requested_outputs(
    result: dict[str, Any],
    *,
    endpoint: str | None,
    request: dict[str, Any] | None,
    normalized_confidence: dict[str, Any] | None = None,
) -> None:
    if endpoint not in {"fold", "fold_all_atom"} or request is None:
        return
    for request_field, response_fields in MANAGED_REQUESTED_OUTPUT_FIELDS.items():
        if request.get(request_field) is not True:
            continue
        missing = [
            field
            for field in response_fields
            if result.get(field) is None
            and (normalized_confidence is None or field not in normalized_confidence)
        ]
        if missing:
            display_names = {"pae": "pAE"}
            raise SchemaDriftError(
                f"managed {endpoint} response is missing requested "
                + ", ".join(display_names.get(field, field) for field in missing),
                raw=result,
            )


def _validate_metric_number(
    value: Any,
    *,
    field: str,
    raw: dict[str, Any],
    low: float | None = None,
    high: float | None = None,
) -> None:
    if isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value):
        raise SchemaDriftError(
            f"managed confidence field {field} must be finite numeric data",
            raw=raw,
        )
    if (low is not None and value < low) or (high is not None and value > high):
        raise SchemaDriftError(
            f"managed confidence field {field} is outside its validated range",
            raw=raw,
        )


def _normalized_plddt_vector(
    value: Any,
    *,
    field: str,
    raw: dict[str, Any],
    expected_residue_length: int | None,
    expected_track_length: int | None,
    sequence_track: str | None,
) -> tuple[list[Any], int | None]:
    expected_lengths = {
        length for length in (expected_residue_length, expected_track_length) if length is not None
    }
    if (
        not isinstance(value, list)
        or not value
        or (expected_lengths and len(value) not in expected_lengths)
    ):
        raise SchemaDriftError(
            f"managed confidence field {field} has an invalid vector shape",
            raw=raw,
        )
    provider_track = (
        sequence_track is not None and "|" in sequence_track and len(value) == len(sequence_track)
    )
    for index, item in enumerate(value):
        is_chain_break = provider_track and sequence_track[index] == "|"
        if is_chain_break:
            if item not in (None, 0, 0.0):
                raise SchemaDriftError(
                    "managed plddt chain-break placeholders must be null or zero",
                    raw=raw,
                )
            continue
        if item is None:
            raise SchemaDriftError(
                "managed plddt confidence cannot be null for an actual residue",
                raw=raw,
            )
        _validate_metric_number(
            item,
            field=field,
            raw=raw,
            low=0,
            high=1,
        )
    if provider_track:
        return (
            [item for item, residue in zip(value, sequence_track, strict=True) if residue != "|"],
            len(value),
        )
    return value, None


def _validate_metric_matrix(
    value: Any,
    *,
    field: str,
    raw: dict[str, Any],
    expected_length: int | None,
    low: float,
    high: float | None,
) -> None:
    if not isinstance(value, list) or not value:
        raise SchemaDriftError(
            f"managed confidence field {field} has an invalid matrix shape",
            raw=raw,
        )
    size = len(value)
    if expected_length is not None and size != expected_length:
        raise SchemaDriftError(
            f"managed confidence field {field} has an invalid matrix shape",
            raw=raw,
        )
    if any(not isinstance(row, list) or len(row) != size for row in value):
        raise SchemaDriftError(
            f"managed confidence field {field} has an invalid matrix shape",
            raw=raw,
        )
    for row in value:
        for item in row:
            if item is None:
                raise SchemaDriftError(
                    f"managed confidence field {field} cannot contain null values",
                    raw=raw,
                )
            _validate_metric_number(
                item,
                field=field,
                raw=raw,
                low=low,
                high=high,
            )


def _plddt_vectors_agree(left: list[Any], right: list[Any]) -> bool:
    return len(left) == len(right) and all(
        math.isclose(
            left_value,
            right_value,
            rel_tol=PLDDT_TRACK_REL_TOLERANCE,
            abs_tol=PLDDT_TRACK_ABS_TOLERANCE,
        )
        for left_value, right_value in zip(left, right, strict=True)
    )


def _all_atom_complex_plddt(
    result: dict[str, Any],
) -> tuple[int | None, list[Any] | None]:
    """Validate the confidence track embedded in the serializable complex state."""

    complex_state = result.get("complex")
    if not isinstance(complex_state, dict):
        return None, None
    sequence = complex_state.get("sequence")
    if not isinstance(sequence, list) or not sequence:
        return None, None
    token_count = len(sequence)
    complex_plddt = complex_state.get("plddt")
    if complex_plddt is None:
        return token_count, None
    if not isinstance(complex_plddt, list) or len(complex_plddt) != token_count:
        raise SchemaDriftError(
            "managed fold_all_atom complex plddt must align with tokens",
            raw=result,
        )
    for item in complex_plddt:
        _validate_metric_number(
            item,
            field="complex.plddt",
            raw=result,
            low=0,
            high=1,
        )
    return token_count, complex_plddt


def _normalized_all_atom_outer_plddt(
    value: Any,
    *,
    field: str,
    raw: dict[str, Any],
    expected_residue_length: int | None,
    expected_track_length: int | None,
    sequence_track: str | None,
    complex_token_count: int | None,
) -> tuple[list[Any], int | None, bool]:
    """Validate an outer pLDDT vector and classify whether its axis is mapped."""

    if (
        isinstance(value, list)
        and sequence_track is not None
        and "|" in sequence_track
        and len(value) == len(sequence_track)
    ):
        normalized, provider_track_length = _normalized_plddt_vector(
            value,
            field=field,
            raw=raw,
            expected_residue_length=expected_residue_length,
            expected_track_length=expected_track_length,
            sequence_track=sequence_track,
        )
    else:
        normalized, provider_track_length = _normalized_plddt_vector(
            value,
            field=field,
            raw=raw,
            expected_residue_length=None,
            expected_track_length=None,
            sequence_track=None,
        )
    mapped = complex_token_count is None or len(normalized) == complex_token_count
    return normalized, provider_track_length, mapped


def _validate_all_atom_plddt_consistency(
    result: dict[str, Any],
    *,
    request: dict[str, Any] | None,
) -> None:
    """Bind provenance confidence to the all-atom track serialized into mmCIF."""

    complex_state = result.get("complex")
    if not isinstance(complex_state, dict):
        return
    complex_plddt = complex_state.get("plddt")
    if not isinstance(complex_plddt, list):
        return

    (
        expected_residue_length,
        expected_track_length,
        sequence_track,
        _,
    ) = _expected_confidence_dimensions(
        result,
        endpoint="fold_all_atom",
        request=request,
    )
    containers = _managed_confidence_containers(result)
    for container in containers:
        value = container.get("plddt")
        if value is None:
            continue
        provider_plddt, _, mapped = _normalized_all_atom_outer_plddt(
            value,
            field="plddt",
            raw=result,
            expected_residue_length=expected_residue_length,
            expected_track_length=expected_track_length,
            sequence_track=sequence_track,
            complex_token_count=len(complex_plddt),
        )
        if mapped and not _plddt_vectors_agree(provider_plddt, complex_plddt):
            raise SchemaDriftError(
                "managed fold_all_atom pLDDT tracks disagree beyond the validated "
                "quantization tolerance",
                raw=result,
            )


def _managed_confidence_contract(
    result: dict[str, Any],
    *,
    endpoint: str | None = None,
    request: dict[str, Any] | None = None,
) -> tuple[dict[str, Any], dict[str, Any]]:
    """Validate confidence once and return its summaries and canonical values."""

    containers = _managed_confidence_containers(result)
    (
        expected_residue_length,
        expected_track_length,
        sequence_track,
        expected_chain_count,
    ) = _expected_confidence_dimensions(
        result,
        endpoint=endpoint,
        request=request,
    )
    metrics: dict[str, Any] = {}
    normalized_values: dict[str, Any] = {}
    plddt_provider_track_length: int | None = None
    metadata: dict[str, dict[str, Any]] = {}
    observed_confidence_fields: set[str] = set()
    mapped_all_atom_fields: set[str] = set()
    unmapped_all_atom_fields: set[str] = set()
    complex_token_count: int | None = None
    complex_plddt: list[Any] | None = None
    if endpoint == "fold_all_atom":
        complex_token_count, complex_plddt = _all_atom_complex_plddt(result)
        if complex_plddt is not None:
            normalized_values["plddt"] = complex_plddt
            metadata["plddt"] = {
                "scale": "0-1",
                "scope": "per-token",
                "index_basis": "returned complex.sequence tokens",
                "source": "complex.plddt",
                "mapping_status": "mapped",
                "completeness": "complete",
            }
            summary = numeric_metric_summary(complex_plddt)
            if isinstance(summary, dict):
                summary.update(metadata["plddt"])
            metrics["plddt"] = summary
            mapped_all_atom_fields.add("complex.plddt")
    all_atom_index_metadata = (
        {"scope": "per-token", "index_basis": "returned complex.sequence tokens"}
        if endpoint == "fold_all_atom"
        else {"scope": "per-residue"}
    )
    all_atom_pair_index_metadata = (
        {"scope": "token-pair", "index_basis": "returned complex.sequence tokens"}
        if endpoint == "fold_all_atom"
        else {"scope": "residue-pair"}
    )
    for container in containers:
        for provider_field, canonical_field in MANAGED_CONFIDENCE_FIELDS:
            value = container.get(provider_field)
            if value is None:
                continue
            observed_confidence_fields.add(canonical_field)
            summary_value = value
            provider_track_length: int | None = None
            metric_field = canonical_field
            if canonical_field == "plddt":
                if endpoint == "fold_all_atom":
                    summary_value, provider_track_length, mapped = (
                        _normalized_all_atom_outer_plddt(
                            value,
                            field=provider_field,
                            raw=result,
                            expected_residue_length=expected_residue_length,
                            expected_track_length=expected_track_length,
                            sequence_track=sequence_track,
                            complex_token_count=complex_token_count,
                        )
                    )
                    if mapped:
                        metadata.setdefault(
                            canonical_field,
                            {
                                "scale": "0-1",
                                **all_atom_index_metadata,
                                "source": f"response.{provider_field}",
                                "mapping_status": "mapped",
                                "completeness": "complete",
                            },
                        )
                        metadata[canonical_field]["outer_mapping_status"] = "mapped"
                        mapped_all_atom_fields.add(provider_field)
                    else:
                        metric_field = "unmapped_plddt"
                        metadata[metric_field] = {
                            "scale": "0-1",
                            "scope": "unmapped",
                            "index_basis": (
                                "provider confidence axis; no mapping to returned "
                                "complex.sequence tokens"
                            ),
                            "source": f"response.{provider_field}",
                            "mapping_status": "unmapped",
                            "completeness": "partial",
                        }
                        unmapped_all_atom_fields.add(provider_field)
                else:
                    summary_value, provider_track_length = _normalized_plddt_vector(
                        value,
                        field=provider_field,
                        raw=result,
                        expected_residue_length=expected_residue_length,
                        expected_track_length=expected_track_length,
                        sequence_track=sequence_track,
                    )
                    metadata[canonical_field] = {
                        "scale": "0-1",
                        **all_atom_index_metadata,
                    }
                if provider_track_length is not None:
                    if (
                        plddt_provider_track_length is not None
                        and plddt_provider_track_length != provider_track_length
                    ):
                        raise SchemaDriftError(
                            "managed plddt provider track shapes disagree",
                            raw=result,
                        )
                    plddt_provider_track_length = provider_track_length
            elif canonical_field in {"ptm", "iptm", "mean_plddt"}:
                _validate_metric_number(
                    value,
                    field=provider_field,
                    raw=result,
                    low=0,
                    high=1,
                )
                metadata[canonical_field] = {"scale": "0-1"}
            elif canonical_field == "pae":
                if endpoint == "fold_all_atom":
                    _validate_metric_matrix(
                        value,
                        field=provider_field,
                        raw=result,
                        expected_length=None,
                        low=0,
                        high=None,
                    )
                    mapped = complex_token_count is None or len(value) == complex_token_count
                    if mapped:
                        metadata[canonical_field] = {
                            "units": "angstrom",
                            **all_atom_pair_index_metadata,
                            "source": f"response.{provider_field}",
                            "mapping_status": "mapped",
                            "completeness": "complete",
                        }
                        mapped_all_atom_fields.add(provider_field)
                    else:
                        metric_field = "unmapped_pae"
                        metadata[metric_field] = {
                            "units": "angstrom",
                            "scope": "unmapped-pair",
                            "index_basis": (
                                "provider confidence axis; no mapping to returned "
                                "complex.sequence tokens"
                            ),
                            "source": f"response.{provider_field}",
                            "mapping_status": "unmapped",
                            "completeness": "partial",
                        }
                        unmapped_all_atom_fields.add(provider_field)
                else:
                    _validate_metric_matrix(
                        value,
                        field=provider_field,
                        raw=result,
                        expected_length=expected_residue_length,
                        low=0,
                        high=None,
                    )
                    metadata[canonical_field] = {
                        "units": "angstrom",
                        **all_atom_pair_index_metadata,
                    }
            elif canonical_field == "pair_chains_iptm":
                _validate_metric_matrix(
                    value,
                    field=provider_field,
                    raw=result,
                    expected_length=expected_chain_count,
                    low=0,
                    high=1,
                )
                metadata[canonical_field] = {"scale": "0-1"}
            if metric_field in normalized_values:
                values_agree = summary_value == normalized_values[metric_field]
                if endpoint == "fold_all_atom" and metric_field == "plddt":
                    values_agree = _plddt_vectors_agree(
                        summary_value,
                        normalized_values[metric_field],
                    )
                if not values_agree:
                    if endpoint == "fold_all_atom" and metric_field == "plddt":
                        raise SchemaDriftError(
                            "managed fold_all_atom pLDDT tracks disagree beyond the "
                            "validated quantization tolerance",
                            raw=result,
                        )
                    raise SchemaDriftError(
                        "managed confidence values disagree for normalized field "
                        f"{metric_field}",
                        raw=result,
                    )
                continue
            normalized_values[metric_field] = summary_value
            summary = numeric_metric_summary(summary_value)
            if isinstance(summary, dict):
                summary.update(metadata[metric_field])
                if provider_track_length is not None:
                    summary["provider_shape"] = [provider_track_length]
                    summary["chain_break_placeholders"] = provider_track_length - len(summary_value)
            metrics[metric_field] = summary
    plddt_summary = metrics.get("plddt")
    if isinstance(plddt_summary, dict) and plddt_provider_track_length is not None:
        plddt_summary["provider_shape"] = [plddt_provider_track_length]
        plddt_summary["chain_break_placeholders"] = plddt_provider_track_length - len(
            normalized_values["plddt"]
        )
    if endpoint in {"fold", "fold_all_atom"} and "plddt" not in metrics:
        raise SchemaDriftError(
            f"managed {endpoint} response is missing required pLDDT confidence",
            raw=result,
        )
    available_confidence = dict(normalized_values)
    available_confidence.update({field: True for field in observed_confidence_fields})
    _require_requested_outputs(
        result,
        endpoint=endpoint,
        request=request,
        normalized_confidence=available_confidence,
    )
    if endpoint == "fold_all_atom":
        metrics["confidence_completeness"] = {
            "status": "partial" if unmapped_all_atom_fields else "complete",
            "index_basis": "returned complex.sequence tokens",
            "mapped_fields": sorted(mapped_all_atom_fields),
            "unmapped_fields": sorted(unmapped_all_atom_fields),
        }
    if metadata:
        metrics["metric_metadata"] = metadata
    return metrics, normalized_values


def managed_confidence_metrics(
    result: dict[str, Any],
    *,
    endpoint: str | None = None,
    request: dict[str, Any] | None = None,
) -> dict[str, Any]:
    """Validate and summarize confidence from a normalized managed response."""

    metrics, _ = _managed_confidence_contract(
        result,
        endpoint=endpoint,
        request=request,
    )
    return metrics


def _load_structure_sdk(endpoint: str):
    """Adapt the readiness-tested endpoint loader to the serializer call sites."""

    components = load_structure_serializer_sdk(endpoint)
    if endpoint == "fold":
        esm_protein, maybe_tensor = components
        return esm_protein, maybe_tensor, None
    (molecular_complex,) = components
    return None, None, molecular_complex


def _expand_compact_sequence_track(
    value: Any,
    *,
    sequence: str,
    field: str,
    chain_break_value: Any,
    raw: dict[str, Any],
) -> Any:
    """Align one residue-only JSON track to the SDK's separator-containing sequence."""

    if value is None or "|" not in sequence:
        return value
    if not isinstance(value, list):
        raise SchemaDriftError(f"managed fold {field} must be a sequence-aligned list", raw=raw)
    residue_count = sum(residue != "|" for residue in sequence)
    if len(value) == len(sequence):
        return value
    if len(value) != residue_count:
        raise SchemaDriftError(
            f"managed fold {field} does not match the submitted sequence length",
            raw=raw,
        )
    iterator = iter(value)
    return [
        deepcopy(chain_break_value) if residue == "|" else next(iterator) for residue in sequence
    ]


def _fold_to_pdb(
    result: dict[str, Any],
    request: dict[str, Any],
    confidence_values: dict[str, Any],
    esm_protein: Any,
    maybe_tensor: Any,
) -> str:
    sequence = request.get("sequence")
    if not isinstance(sequence, str) or not sequence.strip():
        raise ValidationError(
            "managed fold artifact materialization requires the submitted sequence"
        )
    coordinates = result.get("coordinates")
    if not isinstance(coordinates, list) or not coordinates:
        raise SchemaDriftError(
            "managed fold response is missing coordinates",
            raw=result,
        )
    sdk_coordinates = _expand_compact_sequence_track(
        coordinates,
        sequence=sequence,
        field="coordinates",
        chain_break_value=[[None, None, None] for _ in range(37)],
        raw=result,
    )
    sdk_plddt = _expand_compact_sequence_track(
        confidence_values.get("plddt"),
        sequence=sequence,
        field="pLDDT",
        chain_break_value=None,
        raw=result,
    )
    sdk_residue_index = _expand_compact_sequence_track(
        result.get("residue_index"),
        sequence=sequence,
        field="residue_index",
        chain_break_value=None,
        raw=result,
    )
    sdk_entity_id = _expand_compact_sequence_track(
        result.get("entity_id"),
        sequence=sequence,
        field="entity_id",
        chain_break_value=None,
        raw=result,
    )
    protein = esm_protein(
        sequence=sequence,
        coordinates=maybe_tensor(sdk_coordinates, convert_none_to_nan=True),
        ptm=maybe_tensor(confidence_values.get("ptm")),
        plddt=maybe_tensor(sdk_plddt, convert_none_to_nan=True),
        pae=maybe_tensor(confidence_values.get("pae")),
        interface_ptm=maybe_tensor(confidence_values.get("iptm")),
        pair_chains_iptm=maybe_tensor(confidence_values.get("pair_chains_iptm")),
        output_embedding_sequence=maybe_tensor(
            result.get("output_embedding_sequence"),
            convert_none_to_nan=True,
        ),
        output_embedding_pair_pooled=maybe_tensor(
            result.get("output_embedding_pair_pooled"),
            convert_none_to_nan=True,
        ),
        residue_index=maybe_tensor(
            sdk_residue_index,
            convert_none_to_nan=True,
        ),
        entity_id=maybe_tensor(
            sdk_entity_id,
            convert_none_to_nan=True,
        ),
    )
    return protein.to_pdb_string()


def _fold_all_atom_to_mmcif(
    result: dict[str, Any],
    molecular_complex: Any,
) -> str:
    complex_state = result.get("complex")
    if not isinstance(complex_state, dict) or not complex_state:
        raise SchemaDriftError(
            "managed fold_all_atom response is missing the complex state",
            raw=result,
        )
    sdk_state = deepcopy(complex_state)
    metadata = sdk_state.get("metadata")
    if not isinstance(metadata, dict):
        raise SchemaDriftError(
            "managed fold_all_atom complex metadata must be an object",
            raw=result,
        )
    for field in ("entity_lookup", "chain_lookup"):
        lookup = metadata.get(field)
        if not isinstance(lookup, dict):
            raise SchemaDriftError(
                f"managed fold_all_atom complex {field} must be a string lookup",
                raw=result,
            )
        normalized: dict[int, Any] = {}
        original_keys: dict[int, str] = {}
        for key, item in lookup.items():
            if not isinstance(key, str) or not key or any(
                character not in "0123456789" for character in key
            ):
                raise SchemaDriftError(
                    f"managed fold_all_atom complex {field} keys must be nonnegative "
                    "decimal integers",
                    raw=result,
                )
            normalized_key = int(key)
            if normalized_key in normalized:
                raise SchemaDriftError(
                    f"managed fold_all_atom complex {field} keys "
                    f"{original_keys[normalized_key]!r} and {key!r} collide after integer "
                    "normalization",
                    raw=result,
                )
            normalized[normalized_key] = item
            original_keys[normalized_key] = key
        metadata[field] = normalized
    return molecular_complex.from_state_dict(sdk_state).to_mmcif()


def _serialized_atom_record_count(content: str) -> int:
    return sum(
        1
        for line in content.splitlines()
        if line.split(maxsplit=1) and line.split(maxsplit=1)[0] in {"ATOM", "HETATM"}
    )


def _validate_serialized_all_atom_count(
    result: dict[str, Any],
    *,
    observed_count: int,
) -> None:
    """Bind a serialized mmCIF to the validated returned atom array."""

    complex_state = result["complex"]
    expected_count = len(complex_state["atom_positions"])
    if observed_count != expected_count:
        raise SchemaDriftError(
            "managed fold_all_atom mmCIF atom count does not match returned coordinates",
            raw=result,
        )


def _validate_fold_coordinates(
    coordinates: Any,
    *,
    request: dict[str, Any] | None,
    raw: dict[str, Any],
) -> None:
    if not isinstance(coordinates, list) or not coordinates:
        raise SchemaDriftError("managed fold response is missing coordinates", raw=raw)
    sequence_track: list[str | None] = [None] * len(coordinates)
    if request is not None and isinstance((sequence := request.get("sequence")), str):
        biological_track = [residue for residue in sequence if residue != "|"]
        if len(coordinates) == len(sequence):
            sequence_track = list(sequence)
        elif len(coordinates) == len(biological_track):
            sequence_track = biological_track
        else:
            raise SchemaDriftError(
                "managed fold coordinates do not match the submitted sequence length",
                raw=raw,
            )
    for residue_index, (residue, sequence_token) in enumerate(
        zip(coordinates, sequence_track, strict=True)
    ):
        if not isinstance(residue, list) or len(residue) != 37:
            raise SchemaDriftError(
                "managed fold coordinates must use the pinned SDK atom37 shape",
                raw=raw,
            )
        missing_atoms: list[bool] = []
        for atom in residue:
            if not isinstance(atom, list) or len(atom) != 3:
                raise SchemaDriftError(
                    "managed fold coordinates must contain three-dimensional atom positions",
                    raw=raw,
                )
            missing_components = [component is None for component in atom]
            if any(missing_components) and not all(missing_components):
                raise SchemaDriftError(
                    "managed fold atom coordinates must be complete finite triples or all-null",
                    raw=raw,
                )
            missing_atoms.append(all(missing_components))
            for component in atom:
                if component is None:
                    continue
                if (
                    isinstance(component, bool)
                    or not isinstance(component, (int, float))
                    or not math.isfinite(component)
                ):
                    raise SchemaDriftError(
                        "managed fold coordinates must contain finite numbers or nulls",
                        raw=raw,
                    )
        if sequence_token == "|":
            if not all(missing_atoms):
                raise SchemaDriftError(
                    "managed fold chain-break coordinate placeholders must be all-null",
                    raw=raw,
                )
            continue
        if any(missing_atoms[index] for index in (0, 1, 2)):
            raise SchemaDriftError(
                "managed fold biological residues require finite N, CA, and C backbone coordinates "
                f"(missing at residue index {residue_index})",
                raw=raw,
            )


def _validate_string_lookup(value: Any, *, field: str, raw: dict[str, Any]) -> None:
    if not isinstance(value, dict) or any(
        not isinstance(key, str) or not isinstance(item, str) for key, item in value.items()
    ):
        raise SchemaDriftError(
            f"managed fold_all_atom complex {field} must be a string lookup",
            raw=raw,
        )


def _expected_molecular_complex_chains(
    request: dict[str, Any] | None,
) -> list[tuple[str | None, tuple[str, ...], bool]] | None:
    """Return submitted chain identities in the pinned SDK's output representation."""

    if request is None:
        return None
    sequence = request.get("sequence")
    if isinstance(sequence, str) and sequence:
        return [
            (
                None,
                tuple(PROTEIN_RESIDUE_TOKENS.get(residue, "UNK") for residue in chain),
                False,
            )
            for chain in sequence.split("|")
        ]

    all_atom_input = request.get("all_atom_input")
    if not isinstance(all_atom_input, dict):
        return None
    entities = all_atom_input.get("sequences")
    if not isinstance(entities, list) or not entities:
        return None

    chains: list[tuple[str | None, tuple[str, ...], bool]] = []
    for entity in entities:
        if not isinstance(entity, dict):
            return None
        entity_type = entity.get("type")
        if entity_type == "ligand":
            ccd = entity.get("ccd")
            tokens = (ccd[0],) if isinstance(ccd, list) and ccd else ("LIG",)
        elif entity_type in {"protein", "dna", "rna"}:
            entity_sequence = entity.get("sequence")
            if not isinstance(entity_sequence, str) or not entity_sequence:
                return None
            residue_tokens = {
                "protein": PROTEIN_RESIDUE_TOKENS,
                "dna": DNA_RESIDUE_TOKENS,
                "rna": RNA_RESIDUE_TOKENS,
            }[entity_type]
            mapped = [residue_tokens.get(residue, "UNK") for residue in entity_sequence]
            for modification in entity.get("modifications") or []:
                mapped[modification["position"]] = modification["ccd"]
            tokens = tuple(mapped)
        else:
            return None

        identity = entity.get("id")
        chain_names = identity if isinstance(identity, list) else [identity]
        chains.extend((name, tokens, entity_type == "ligand") for name in chain_names)
    return chains


def _bind_molecular_complex_chains(
    value: dict[str, Any],
    *,
    request: dict[str, Any] | None,
    raw: dict[str, Any],
) -> None:
    """Bind returned residue and chain identities to the normalized submitted request."""

    expected = _expected_molecular_complex_chains(request)
    if expected is None:
        return
    chain_ids = value.get("chain_id")
    if chain_ids is None:
        if len(expected) != 1:
            raise SchemaDriftError(
                "managed fold_all_atom multichain response is missing chain_id",
                raw=raw,
            )
        chain_ids = [0] * len(value["sequence"])

    lookup = value["metadata"]["chain_lookup"]
    if len(set(lookup.values())) != len(lookup):
        raise SchemaDriftError(
            "managed fold_all_atom complex chain_lookup repeats a chain identity",
            raw=raw,
        )
    observed: dict[str, list[str]] = {}
    observed_hetero: dict[str, set[bool]] = {}
    hetero = value.get("atom_hetero")
    for token, chain_index, atom_bounds in zip(
        value["sequence"], chain_ids, value["token_to_atoms"], strict=True
    ):
        name = lookup.get(str(chain_index))
        if name is None:
            raise SchemaDriftError(
                "managed fold_all_atom complex chain_id does not resolve in chain_lookup",
                raw=raw,
            )
        observed.setdefault(name, []).append(token)
        if hetero is not None:
            start, end = atom_bounds
            observed_hetero.setdefault(name, set()).update(hetero[start:end])

    if set(observed) != set(lookup.values()):
        raise SchemaDriftError(
            "managed fold_all_atom complex chain_lookup contains an unrepresented chain",
            raw=raw,
        )
    named = {name: (tokens, is_ligand) for name, tokens, is_ligand in expected if name is not None}
    anonymous = [(tokens, is_ligand) for name, tokens, is_ligand in expected if name is None]
    if len(named) + len(anonymous) != len(expected) or len(observed) != len(expected):
        raise SchemaDriftError(
            "managed fold_all_atom response chain composition does not match the request",
            raw=raw,
        )
    if not set(named).issubset(observed):
        raise SchemaDriftError(
            "managed fold_all_atom response chain identities do not match the request",
            raw=raw,
        )
    for name, (tokens, is_ligand) in named.items():
        if tuple(observed[name]) != tokens:
            raise SchemaDriftError(
                f"managed fold_all_atom response residue identities do not match requested chain {name}",
                raw=raw,
            )
        if hetero is not None and observed_hetero.get(name) != {is_ligand}:
            raise SchemaDriftError(
                f"managed fold_all_atom response entity type does not match requested chain {name}",
                raw=raw,
            )

    unmatched: list[tuple[tuple[str, ...], bool | None]] = []
    for name, tokens in observed.items():
        if name in named:
            continue
        if hetero is None:
            unmatched.append((tuple(tokens), None))
            continue
        identities = observed_hetero.get(name, set())
        if len(identities) != 1:
            raise SchemaDriftError(
                "managed fold_all_atom response chain mixes polymer and ligand atoms",
                raw=raw,
            )
        unmatched.append((tuple(tokens), next(iter(identities))))
    if hetero is None:
        expected_anonymous = Counter(tokens for tokens, _ in anonymous)
        actual_anonymous = Counter(tokens for tokens, _ in unmatched)
    else:
        expected_anonymous = Counter(anonymous)
        actual_anonymous = Counter(unmatched)
    if expected_anonymous != actual_anonymous:
        raise SchemaDriftError(
            "managed fold_all_atom response residue or entity identities do not match the request",
            raw=raw,
        )


def _validate_molecular_complex_state(
    value: Any,
    *,
    raw: dict[str, Any],
    request: dict[str, Any] | None,
) -> None:
    if not isinstance(value, dict) or not value:
        raise SchemaDriftError(
            "managed fold_all_atom response is missing the complex state",
            raw=raw,
        )
    required = {
        "id",
        "sequence",
        "atom_positions",
        "atom_elements",
        "token_to_atoms",
        "plddt",
        "metadata",
    }
    allowed = required | {"chain_id", "atom_names", "atom_hetero"}
    missing = sorted(required - value.keys())
    unknown = sorted(value.keys() - allowed)
    if missing or unknown:
        detail = []
        if missing:
            detail.append("missing " + ", ".join(missing))
        if unknown:
            detail.append("unknown " + ", ".join(unknown))
        raise SchemaDriftError(
            "managed fold_all_atom complex state has " + "; ".join(detail),
            raw=raw,
        )
    if not isinstance(value["id"], str) or not value["id"]:
        raise SchemaDriftError(
            "managed fold_all_atom complex id must be a nonempty string",
            raw=raw,
        )
    sequence = value["sequence"]
    if (
        not isinstance(sequence, list)
        or not sequence
        or any(not isinstance(token, str) or not token for token in sequence)
    ):
        raise SchemaDriftError(
            "managed fold_all_atom complex sequence must contain token strings",
            raw=raw,
        )
    atom_positions = value["atom_positions"]
    if not isinstance(atom_positions, list) or not atom_positions:
        raise SchemaDriftError(
            "managed fold_all_atom complex atom_positions must be nonempty",
            raw=raw,
        )
    for position in atom_positions:
        if (
            not isinstance(position, list)
            or len(position) != 3
            or any(
                isinstance(component, bool)
                or not isinstance(component, (int, float))
                or not math.isfinite(component)
                for component in position
            )
        ):
            raise SchemaDriftError(
                "managed fold_all_atom complex atom_positions must be finite 3D coordinates",
                raw=raw,
            )
    atom_count = len(atom_positions)
    atom_elements = value["atom_elements"]
    if (
        not isinstance(atom_elements, list)
        or len(atom_elements) != atom_count
        or any(not isinstance(element, str) or not element for element in atom_elements)
    ):
        raise SchemaDriftError(
            "managed fold_all_atom complex atom_elements must align with atoms",
            raw=raw,
        )
    token_to_atoms = value["token_to_atoms"]
    if not isinstance(token_to_atoms, list) or len(token_to_atoms) != len(sequence):
        raise SchemaDriftError(
            "managed fold_all_atom complex token_to_atoms must align with tokens",
            raw=raw,
        )
    expected_start = 0
    for bounds in token_to_atoms:
        if (
            not isinstance(bounds, list)
            or len(bounds) != 2
            or any(isinstance(item, bool) or not isinstance(item, int) for item in bounds)
        ):
            raise SchemaDriftError(
                "managed fold_all_atom complex token_to_atoms bounds are invalid",
                raw=raw,
            )
        start, end = bounds
        if start != expected_start or not start < end or end > atom_count:
            raise SchemaDriftError(
                "managed fold_all_atom complex token_to_atoms must partition the atom array",
                raw=raw,
            )
        expected_start = end
    if expected_start != atom_count:
        raise SchemaDriftError(
            "managed fold_all_atom complex token_to_atoms must cover every atom",
            raw=raw,
        )
    chain_id = value.get("chain_id")
    if chain_id is not None and (
        not isinstance(chain_id, list)
        or len(chain_id) != len(sequence)
        or any(isinstance(item, bool) or not isinstance(item, int) for item in chain_id)
    ):
        raise SchemaDriftError(
            "managed fold_all_atom complex chain_id must align with tokens",
            raw=raw,
        )
    inner_plddt = value["plddt"]
    if not isinstance(inner_plddt, list) or len(inner_plddt) != len(sequence):
        raise SchemaDriftError(
            "managed fold_all_atom complex plddt must align with tokens",
            raw=raw,
        )
    for item in inner_plddt:
        _validate_metric_number(item, field="complex.plddt", raw=raw, low=0, high=1)
    metadata = value["metadata"]
    if not isinstance(metadata, dict):
        raise SchemaDriftError(
            "managed fold_all_atom complex metadata must be an object",
            raw=raw,
        )
    metadata_allowed = {"entity_lookup", "chain_lookup", "assembly_composition"}
    if set(metadata) - metadata_allowed or not {
        "entity_lookup",
        "chain_lookup",
    }.issubset(metadata):
        raise SchemaDriftError(
            "managed fold_all_atom complex metadata fields are invalid",
            raw=raw,
        )
    _validate_string_lookup(metadata["entity_lookup"], field="entity_lookup", raw=raw)
    _validate_string_lookup(metadata["chain_lookup"], field="chain_lookup", raw=raw)
    assembly = metadata.get("assembly_composition")
    if assembly is not None and (
        not isinstance(assembly, dict)
        or any(
            not isinstance(key, str)
            or not isinstance(items, list)
            or any(not isinstance(item, str) for item in items)
            for key, items in assembly.items()
        )
    ):
        raise SchemaDriftError(
            "managed fold_all_atom complex assembly_composition is invalid",
            raw=raw,
        )
    atom_names = value.get("atom_names")
    if atom_names is not None and (
        not isinstance(atom_names, list)
        or len(atom_names) != atom_count
        or any(not isinstance(name, str) or not name for name in atom_names)
    ):
        raise SchemaDriftError(
            "managed fold_all_atom complex atom_names must align with atoms",
            raw=raw,
        )
    atom_hetero = value.get("atom_hetero")
    if atom_hetero is not None and (
        not isinstance(atom_hetero, list)
        or len(atom_hetero) != atom_count
        or any(not isinstance(item, bool) for item in atom_hetero)
    ):
        raise SchemaDriftError(
            "managed fold_all_atom complex atom_hetero must align with atoms",
            raw=raw,
        )
    _bind_molecular_complex_chains(value, request=request, raw=raw)


def validate_managed_structure_response(
    endpoint: str,
    result: dict[str, Any],
    request: dict[str, Any] | None = None,
) -> None:
    """Validate fold structures against the pinned SDK's JSON state contract."""

    if endpoint not in {"fold", "fold_all_atom"}:
        return
    artifact_key = "structure_artifact"
    if artifact_key in result:
        raise SchemaDriftError(
            f"managed {endpoint} response contains reserved {artifact_key}",
            raw=result,
        )
    if endpoint == "fold":
        _validate_fold_coordinates(result.get("coordinates"), request=request, raw=result)
        return
    _validate_molecular_complex_state(result.get("complex"), raw=result, request=request)
    _validate_all_atom_plddt_consistency(result, request=request)


def managed_structure_quality_warnings(
    endpoint: str,
    result: dict[str, Any],
    request: dict[str, Any],
) -> list[dict[str, Any]]:
    """Return conservative, non-mutating warnings for declared covalent geometry.

    This check runs only after the all-atom response contract has validated. It
    deliberately does not infer undeclared bonds, repair coordinates, or treat
    predicted confidence as evidence that the chemistry is sound.
    """

    if endpoint != "fold_all_atom":
        return []
    all_atom_input = request.get("all_atom_input")
    if not isinstance(all_atom_input, dict):
        return []
    bonds = all_atom_input.get("covalent_bonds")
    if not isinstance(bonds, list) or not bonds:
        return []
    complex_state = result.get("complex")
    if not isinstance(complex_state, dict):
        return []

    sequence = complex_state["sequence"]
    chain_ids = complex_state.get("chain_id")
    if chain_ids is None:
        chain_ids = [0] * len(sequence)
    chain_lookup = complex_state["metadata"]["chain_lookup"]
    tokens_by_chain: dict[str, list[int]] = {}
    for token_index, chain_index in enumerate(chain_ids):
        chain_name = chain_lookup.get(str(chain_index))
        if chain_name is not None:
            tokens_by_chain.setdefault(chain_name, []).append(token_index)

    atom_positions = complex_state["atom_positions"]
    atom_elements = complex_state["atom_elements"]
    atom_names = complex_state.get("atom_names")
    token_to_atoms = complex_state["token_to_atoms"]

    def resolve_endpoint(
        bond: dict[str, Any],
        suffix: str,
    ) -> tuple[dict[str, Any], list[float] | None]:
        chain_name = bond[f"chain_id{suffix}"]
        residue_index = bond[f"res_idx{suffix}"]
        atom_index = bond[f"atom_idx{suffix}"]
        description: dict[str, Any] = {
            "chain_id": chain_name,
            "residue_index_zero_based": residue_index,
            "atom_index_zero_based": atom_index,
        }
        chain_tokens = tokens_by_chain.get(chain_name, [])
        if residue_index >= len(chain_tokens):
            return description, None
        token_index = chain_tokens[residue_index]
        start, end = token_to_atoms[token_index]
        absolute_atom_index = start + atom_index
        if absolute_atom_index >= end:
            return description, None
        description["element"] = atom_elements[absolute_atom_index]
        if isinstance(atom_names, list):
            description["atom_name"] = atom_names[absolute_atom_index]
        return description, atom_positions[absolute_atom_index]

    warnings: list[dict[str, Any]] = []
    for bond_index, item in enumerate(bonds):
        if not isinstance(item, dict):
            continue
        endpoint_1, position_1 = resolve_endpoint(item, "1")
        endpoint_2, position_2 = resolve_endpoint(item, "2")
        endpoints = [endpoint_1, endpoint_2]
        if position_1 is None or position_2 is None:
            warnings.append(
                {
                    "code": "covalent-attachment-atom-unresolved",
                    "severity": "warning",
                    "category": "chemistry",
                    "bond_index_zero_based": bond_index,
                    "endpoints": endpoints,
                    "message": (
                        "A requested covalent attachment atom is absent from the returned "
                        "atom map; verify the zero-based indices against the molecular graph."
                    ),
                    "coordinates_modified": False,
                }
            )
            continue
        distance = math.dist(position_1, position_2)
        if (
            distance >= COVALENT_BOND_SEVERE_MIN_ANGSTROM
            and distance <= COVALENT_BOND_SEVERE_MAX_ANGSTROM
        ):
            continue
        warnings.append(
            {
                "code": "implausible-covalent-bond-length",
                "severity": "warning",
                "category": "severe-geometry",
                "bond_index_zero_based": bond_index,
                "endpoints": endpoints,
                "distance_angstrom": round(distance, 3),
                "conservative_expected_range_angstrom": {
                    "minimum": COVALENT_BOND_SEVERE_MIN_ANGSTROM,
                    "maximum": COVALENT_BOND_SEVERE_MAX_ANGSTROM,
                },
                "message": (
                    "The predicted attachment distance is outside a conservative generic "
                    "covalent-bond range. Verify the chemistry; predicted confidence is not "
                    "a chemistry check."
                ),
                "coordinates_modified": False,
            }
        )
    return warnings


def _normalize_all_atom_confidence_output(
    result: dict[str, Any],
    confidence_metrics: dict[str, Any],
    confidence_values: dict[str, Any],
) -> dict[str, Any]:
    """Separate returned-token confidence from provider-only unmapped axes."""

    normalized = deepcopy(result)
    metadata = confidence_metrics.get("metric_metadata")
    if not isinstance(metadata, dict):
        metadata = {}

    # Remove provider arrays from every supported confidence envelope before
    # exposing a single canonical mapped value or an explicit quarantine.
    for container in _managed_confidence_containers(normalized):
        container.pop("plddt", None)
        container.pop("pae", None)

    mapped_plddt = confidence_values.get("plddt")
    if mapped_plddt is not None:
        normalized["plddt"] = deepcopy(mapped_plddt)
    mapped_pae = confidence_values.get("pae")
    if mapped_pae is not None:
        normalized["pae"] = deepcopy(mapped_pae)

    quarantine: dict[str, Any] = {}
    for canonical_name in ("plddt", "pae"):
        metric_name = f"unmapped_{canonical_name}"
        values = confidence_values.get(metric_name)
        metric_metadata = metadata.get(metric_name)
        if values is None or not isinstance(metric_metadata, dict):
            continue
        quarantine[canonical_name] = {
            "values": deepcopy(values),
            "metadata": deepcopy(metric_metadata),
        }
    if quarantine:
        normalized["unmapped_confidence"] = quarantine

    completeness = confidence_metrics.get("confidence_completeness")
    normalized["confidence_contract"] = {
        "completeness": deepcopy(completeness),
        "metric_metadata": deepcopy(metadata),
    }
    return normalized


def materialize_managed_structure(
    endpoint: str,
    result: dict[str, Any],
    request: dict[str, Any],
    output_dir: Path,
    *,
    artifact_publisher: Callable[[Path, bytes, str], dict[str, Any]] | None = None,
) -> tuple[dict[str, Any], list[dict[str, Any]], str]:
    """Serialize a real managed fold response with the exact pinned ESM SDK."""

    if endpoint not in {"fold", "fold_all_atom"}:
        raise ValidationError("managed structure materialization requires a fold endpoint")

    artifact_key = "structure_artifact"
    for reserved in (
        artifact_key,
        "confidence_contract",
        "unmapped_confidence",
        "quality_warnings",
        "presentation_request",
    ):
        if reserved in result:
            raise SchemaDriftError(
                f"managed {endpoint} response contains reserved {reserved}",
                raw=result,
            )
    validate_managed_structure_response(endpoint, result, request)

    if endpoint == "fold":
        sequence = request.get("sequence")
        if not isinstance(sequence, str) or not sequence.strip():
            raise ValidationError(
                "managed fold artifact materialization requires the submitted sequence"
            )
        confidence_metrics, confidence_values = _managed_confidence_contract(
            result,
            endpoint=endpoint,
            request=request,
        )
    else:
        confidence_metrics, confidence_values = _managed_confidence_contract(
            result,
            endpoint=endpoint,
            request=request,
        )

    esm_revision = verify_installed_vcs_revision("esm", ESM_GIT_REVISION)
    try:
        esm_protein, maybe_tensor, molecular_complex = _load_structure_sdk(endpoint)
    except ImportError as exc:
        raise ValidationError("pinned Biohub esm SDK is not installed") from exc

    try:
        if endpoint == "fold":
            content = _fold_to_pdb(
                result,
                request,
                confidence_values,
                esm_protein,
                maybe_tensor,
            )
            path = output_dir / "prediction.pdb"
            media_type = "chemical/x-pdb"
        else:
            content = _fold_all_atom_to_mmcif(result, molecular_complex)
            path = output_dir / "prediction.cif"
            media_type = "chemical/x-mmcif"
    except SchemaDriftError:
        raise
    except (
        AssertionError,
        AttributeError,
        IndexError,
        KeyError,
        RuntimeError,
        TypeError,
        ValueError,
    ) as exc:
        raise SchemaDriftError(
            f"managed {endpoint} response could not be serialized by the pinned SDK",
            raw=result,
        ) from exc

    if not isinstance(content, str) or not content.strip():
        raise SchemaDriftError(
            f"managed {endpoint} serialization produced an empty structure",
            raw=result,
        )
    atom_record_count = _serialized_atom_record_count(content)
    if atom_record_count == 0:
        raise SchemaDriftError(
            f"managed {endpoint} serialization produced no atom records",
            raw=result,
        )
    if endpoint == "fold_all_atom":
        _validate_serialized_all_atom_count(
            result,
            observed_count=atom_record_count,
        )

    encoded = content.encode("utf-8")
    if artifact_publisher is None:
        write_bytes_atomic(path, encoded)
        record = artifact_record(path, media_type=media_type)
    else:
        record = artifact_publisher(path, encoded, media_type)
    record["evidence_class"] = "model_hypothesis"
    normalized = (
        _normalize_all_atom_confidence_output(
            result,
            confidence_metrics,
            confidence_values,
        )
        if endpoint == "fold_all_atom"
        else dict(result)
    )
    normalized.pop("coordinates" if endpoint == "fold" else "complex", None)
    normalized[artifact_key] = record
    return normalized, [record], esm_revision

SHA-256: cfb28968202c8790ca25a4a6c322c16bc392629397feb3bccb4b10c06bb890f7