Back to home page

EIC code displayed by LXR

 
 

    


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

0001 from types import MappingProxyType
0002 from typing import Any, Mapping
0003 
0004 from .catalog import DatasetCatalog
0005 from .datasets import DatasetDID
0006 from .models import Decision, DecisionContext, FileDID, Site
0007 from .policy import DecisionPolicy
0008 
0009 
0010 class DecisionBox:
0011     """Apply ePIC processing decisions through site-specific datasets."""
0012 
0013     def __init__(
0014         self,
0015         catalog: DatasetCatalog,
0016         sites: tuple[Site, ...] = (Site("E1_BNL"), Site("E1_JLAB")),
0017         manage_full_dataset: bool = True,
0018         site_dataset_template: str | None = None,
0019     ):
0020         if not sites:
0021             raise ValueError("at least one site is required")
0022         self.catalog = catalog
0023         self.sites = sites
0024         self.manage_full_dataset = manage_full_dataset
0025         self.site_dataset_template = site_dataset_template
0026         self._run_sequences: dict[str, int] = {}
0027 
0028     def create_run(self, run_dataset: str) -> tuple[str, ...]:
0029         run_did = DatasetDID.parse(run_dataset)
0030         dataset_dids = []
0031         if self.manage_full_dataset:
0032             dataset_dids.append(str(run_did))
0033         dataset_dids.extend(str(self._site_dataset(run_did, site.name)) for site in self.sites)
0034         for dataset_did in dataset_dids:
0035             self.catalog.ensure_dataset(dataset_did, open_dataset=True)
0036         return tuple(dataset_dids)
0037 
0038     def decide_file(
0039         self,
0040         run_dataset: str,
0041         file_did: FileDID,
0042         policy: DecisionPolicy,
0043         *,
0044         message: Mapping[str, Any] | None = None,
0045         run_conditions: Mapping[str, Any] | None = None,
0046         metadata: Mapping[str, Any] | None = None,
0047     ) -> Decision:
0048         run_did = DatasetDID.parse(run_dataset)
0049         self.create_run(run_dataset)
0050 
0051         message_data = dict(message or {})
0052         sequence = self._sequence_from_message(message_data, run_dataset)
0053         context = DecisionContext(
0054             run_dataset=run_dataset,
0055             file_did=file_did,
0056             available_sites=self.sites,
0057             sequence=sequence,
0058             run_number=run_did.run_number,
0059             message=MappingProxyType(message_data),
0060             run_conditions=MappingProxyType(dict(run_conditions or {})),
0061             metadata=MappingProxyType(dict(metadata or {})),
0062         )
0063         assignment = policy.choose_sites(context)
0064         self._run_sequences[run_dataset] = sequence + 1
0065 
0066         site_datasets = [
0067             str(self._site_dataset(run_did, site.name))
0068             for site in assignment.sites
0069         ]
0070         for site_dataset in site_datasets:
0071             self.catalog.ensure_dataset(site_dataset, open_dataset=True)
0072 
0073         if self.manage_full_dataset:
0074             self.catalog.attach_file(str(run_did), file_did)
0075         for site_dataset in site_datasets:
0076             self.catalog.attach_file(site_dataset, file_did)
0077 
0078         return Decision(
0079             run_dataset=run_dataset,
0080             full_dataset=str(run_did),
0081             file_did=file_did,
0082             site_datasets=tuple(site_datasets),
0083             reason=assignment.reason,
0084         )
0085 
0086     def close_run(self, run_dataset: str) -> tuple[str, ...]:
0087         run_did = DatasetDID.parse(run_dataset)
0088         dataset_dids = []
0089         if self.manage_full_dataset:
0090             dataset_dids.append(str(run_did))
0091         dataset_dids.extend(str(self._site_dataset(run_did, site.name)) for site in self.sites)
0092         for dataset_did in dataset_dids:
0093             self.catalog.close_dataset(dataset_did)
0094         return tuple(dataset_dids)
0095 
0096     def _sequence_for(self, run_dataset: str) -> int:
0097         remembered_sequence = self._run_sequences.get(run_dataset, 0)
0098         if not self.manage_full_dataset:
0099             return remembered_sequence
0100         snapshot = getattr(self.catalog, "snapshot", lambda: {"datasets": {}})()
0101         files = snapshot.get("datasets", {}).get(run_dataset, {}).get("files", [])
0102         return max(remembered_sequence, len(files))
0103 
0104     def _sequence_from_message(self, message: Mapping[str, Any], run_dataset: str) -> int:
0105         for key in ("decision_sequence", "sequence"):
0106             try:
0107                 return int(message[key])
0108             except (KeyError, TypeError, ValueError):
0109                 continue
0110         return self._sequence_for(run_dataset)
0111 
0112     def _site_dataset(self, run_did: DatasetDID, site_name: str) -> DatasetDID:
0113         return run_did.site_dataset(site_name, self.site_dataset_template)