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)