from __future__ import annotations from abc import ABC, abstractmethod from typing import Any from staplebridge.chemistry.state import StapleState from staplebridge.data.schemas import BuildingBlock class PeptidePriorBase(ABC): @abstractmethod def score_transition(self, old_state: StapleState, new_state: StapleState, context: dict[str, Any] | None = None) -> float: pass def batch_score_transitions( self, old_state: StapleState, new_states: list[StapleState], context: dict[str, Any] | None = None, ) -> list[float]: """Score N candidate transitions from ``old_state`` at once. Default implementation just loops ``score_transition``; heavyweight priors (ESM2) override this so they can share one model forward across all candidates. Non-sequence-changing candidates are expected to return exactly 0.0. """ return [ self.score_transition(old_state, new, context) for new in new_states ] def prewarm_requests( self, pairs: list[tuple[StapleState, list[StapleState]]] ) -> None: """Prefetch model outputs for many (z, candidates) pairs at once. Default is a noop — heavy priors (ESM2) override this to run one batched model forward covering every request across all pairs, so a subsequent per-pair ``batch_score_transitions`` call becomes a pure cache-lookup. """ del pairs class AnchorPriorBase(ABC): @abstractmethod def score_anchor(self, sequence: list[str], anchor_pair: tuple[int, int] | None, context: dict[str, Any] | None = None) -> float: pass class BlockPriorBase(ABC): @abstractmethod def score_block( self, sequence: list[str], anchor_pair: tuple[int, int] | None, block: BuildingBlock | None, context: dict[str, Any] | None = None, ) -> float: pass class GeometryOracleBase(ABC): @abstractmethod def ctype( self, sequence: list[str], anchor_pair: tuple[int, int] | None, block: BuildingBlock | None, *, peptide_ca: list[tuple[float, float, float]] | None = None, ) -> bool: pass @abstractmethod def cgeom( self, sequence: list[str], anchor_pair: tuple[int, int] | None, block: BuildingBlock | None, *, peptide_ca: list[tuple[float, float, float]] | None = None, ) -> float: pass