Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-09-01 09:34:30

0001 import hashlib
0002 from abc import ABC, abstractmethod
0003 
0004 from .models import DecisionContext, FileDID, Site, SiteAssignment
0005 
0006 
0007 class DecisionPolicy(ABC):
0008     """Policy deciding which site datasets receive a file DID."""
0009 
0010     def choose_sites(
0011         self,
0012         context: DecisionContext | FileDID,
0013         sites: tuple[Site, ...] | None = None,
0014         sequence: int | None = None,
0015     ) -> SiteAssignment:
0016         """Return the sites that should process the context's file DID."""
0017         if not isinstance(context, DecisionContext):
0018             if sites is None or sequence is None:
0019                 raise TypeError("choose_sites requires a DecisionContext")
0020             context = DecisionContext(
0021                 run_dataset="",
0022                 file_did=context,
0023                 available_sites=sites,
0024                 sequence=sequence,
0025             )
0026         return self._choose_sites(context)
0027 
0028     @abstractmethod
0029     def _choose_sites(self, context: DecisionContext) -> SiteAssignment:
0030         """Policy implementation for a normalized decision context."""
0031 
0032 
0033 class RoundRobinPolicy(DecisionPolicy):
0034     def _choose_sites(self, context: DecisionContext) -> SiteAssignment:
0035         sites = context.available_sites
0036         _require_sites(sites)
0037         site = sites[context.sequence % len(sites)]
0038         return SiteAssignment.from_sites(context.file_did, [site], f"round-robin sequence={context.sequence}")
0039 
0040 
0041 class HashPolicy(DecisionPolicy):
0042     def _choose_sites(self, context: DecisionContext) -> SiteAssignment:
0043         sites = context.available_sites
0044         _require_sites(sites)
0045         digest = hashlib.sha256(str(context.file_did).encode("utf-8")).digest()
0046         idx = int.from_bytes(digest[:8], "big") % len(sites)
0047         return SiteAssignment.from_sites(context.file_did, [sites[idx]], "sha256 modulo site count")
0048 
0049 
0050 class AllSitesPolicy(DecisionPolicy):
0051     def _choose_sites(self, context: DecisionContext) -> SiteAssignment:
0052         sites = context.available_sites
0053         _require_sites(sites)
0054         return SiteAssignment.from_sites(context.file_did, sites, "all sites selected")
0055 
0056 
0057 class NoSitesPolicy(DecisionPolicy):
0058     def _choose_sites(self, context: DecisionContext) -> SiteAssignment:
0059         return SiteAssignment.from_sites(context.file_did, [], "no site selected")
0060 
0061 
0062 class ExplicitPolicy(DecisionPolicy):
0063     def __init__(self, selected_sites: tuple[Site, ...]):
0064         self.selected_sites = selected_sites
0065 
0066     def _choose_sites(self, context: DecisionContext) -> SiteAssignment:
0067         sites = context.available_sites
0068         allowed = {site.name: site for site in sites}
0069         unknown = [site.name for site in self.selected_sites if site.name not in allowed]
0070         if unknown:
0071             raise ValueError(f"unknown site(s): {', '.join(unknown)}")
0072         selected = [allowed[site.name] for site in self.selected_sites]
0073         return SiteAssignment.from_sites(context.file_did, selected, "explicit site list")
0074 
0075 
0076 def build_policy(name: str, selected_sites: tuple[Site, ...] = ()) -> DecisionPolicy:
0077     normalized = name.strip().lower()
0078     if normalized == "round-robin":
0079         return RoundRobinPolicy()
0080     if normalized == "hash":
0081         return HashPolicy()
0082     if normalized in {"both", "all", "broadcast"}:
0083         return AllSitesPolicy()
0084     if normalized in {"none", "skip"}:
0085         return NoSitesPolicy()
0086     if normalized == "explicit":
0087         return ExplicitPolicy(selected_sites)
0088     raise ValueError(f"unknown policy {name!r}")
0089 
0090 
0091 def _require_sites(sites: tuple[Site, ...]) -> None:
0092     if not sites:
0093         raise ValueError("at least one site is required")