"""Schema describing data acquisition metadata and configurations""" import logging import re from datetime import datetime, timedelta, timezone from decimal import Decimal from typing import Annotated, List, Literal, Optional, Union from zoneinfo import ZoneInfo from aind_data_schema_models.modalities import Modality from aind_data_schema_models.stimulus_modality import StimulusModality from aind_data_schema_models.units import MassUnit, VolumeUnit from pydantic import Field, SkipValidation, field_validator, model_validator from pydantic_extra_types.timezone_name import TimeZoneName from aind_data_schema.base import AwareDatetimeWithDefault, DataCoreModel, DataModel, DiscriminatedList, GenericModel from aind_data_schema.components.configs import ( AirPuffConfig, CatheterConfig, DetectorConfig, EphysAssemblyConfig, FiberAssemblyConfig, ImagingConfig, LaserConfig, LickSpoutConfig, LightEmittingDiodeConfig, ManipulatorConfig, MISCameraConfig, MousePlatformConfig, MRIScan, OlfactometerConfig, PatchCordConfig, ProbeConfig, SampleChamberConfig, SpeakerConfig, ) from aind_data_schema.components.connections import Connection from aind_data_schema.components.coordinates import CoordinateSystem from aind_data_schema.components.identifiers import Code, ProtocolListMixin from aind_data_schema.components.measurements import CALIBRATIONS, Maintenance from aind_data_schema.components.reagent import Reagent from aind_data_schema.components.subject_procedures import BrainInjection, Injection from aind_data_schema.components.surgery_procedures import Anaesthetic from aind_data_schema.utils.merge import ( merge_coordinate_systems, merge_notes, merge_optional_list, merge_str_alphabetical, remove_duplicates, ) from aind_data_schema.utils.validators import ( TimeValidation, extract_timezone_from_datetime, subject_specimen_id_compatibility, ) logger = logging.getLogger(__name__) # Define the requirements for each modality # Define the mapping of modalities to their required device types # The list of list pattern is used to allow for multiple options within a group, so e.g. # FIB requires a light config (one of the options) plus a fiber connection config and a fiber module CONFIG_REQUIREMENTS = { Modality.ECEPHYS.abbreviation: [[EphysAssemblyConfig, ProbeConfig, ManipulatorConfig]], Modality.FIB.abbreviation: [[LightEmittingDiodeConfig, LaserConfig], [PatchCordConfig, FiberAssemblyConfig]], Modality.POPHYS.abbreviation: [[ImagingConfig]], Modality.MRI.abbreviation: [[MRIScan]], Modality.SPIM.abbreviation: [[ImagingConfig], [SampleChamberConfig]], Modality.SLAP2.abbreviation: [[ImagingConfig]], } SPECIMEN_MODALITIES = [Modality.SPIM.abbreviation, Modality.CONFOCAL.abbreviation] class AcquisitionSubjectDetails(DataModel): """Details about the subject during an acquisition""" animal_weight_prior: Optional[Decimal] = Field( default=None, title="Animal weight (g)", description="Animal weight before procedure", ) animal_weight_post: Optional[Decimal] = Field( default=None, title="Animal weight (g)", description="Animal weight after procedure", ) weight_unit: MassUnit = Field(default=MassUnit.G, title="Weight unit") anaesthesia: Optional[Anaesthetic] = Field( default=None, title="Anaesthesia", description=("Anaesthesia present during entire acquisition, use Manipulation for partial anaesthesia"), ) mouse_platform_name: str = Field( ..., title="Mouse platform", description="The surface that the mouse is on during the acquisition" ) reward_consumed_total: Optional[Decimal] = Field(default=None, title="Total reward consumed (mL)") reward_consumed_unit: Optional[VolumeUnit] = Field(default=None, title="Reward consumed unit") class PerformanceMetrics(DataModel): """Summary of a StimulusEpoch""" output_parameters: Optional[GenericModel] = Field(default=None, title="Additional metrics") reward_consumed_during_epoch: Optional[Decimal] = Field(default=None, title="Reward consumed during training (uL)") reward_consumed_unit: Optional[VolumeUnit] = Field(default=None, title="Reward consumed unit") trials_total: Optional[int] = Field(default=None, title="Total trials") trials_finished: Optional[int] = Field(default=None, title="Finished trials") trials_rewarded: Optional[int] = Field(default=None, title="Rewarded trials") class DataStream(DataModel): """A set of devices that are acquiring data and their configurations starting and stopping at approximately the same time. """ object_type: Literal["DataStream"] = "DataStream" stream_start_time: Annotated[ AwareDatetimeWithDefault, Field(..., title="Stream start time"), TimeValidation.BETWEEN, ] stream_end_time: Annotated[ AwareDatetimeWithDefault, Field(..., title="Stream stop time"), TimeValidation.BETWEEN, ] modalities: List[Modality.ONE_OF] = Field( ..., title="Modalities", description="Modalities that are acquired in this stream" ) code: Optional[List[Code]] = Field(default=None, title="Acquisition code") notes: Optional[str] = Field(default=None, title="Notes") active_devices: List[str] = Field( ..., title="Active devices", description="Device names must match devices in the Instrument", ) configurations: DiscriminatedList[ LightEmittingDiodeConfig | LaserConfig | ManipulatorConfig | DetectorConfig | PatchCordConfig | FiberAssemblyConfig | MISCameraConfig | MRIScan | LickSpoutConfig | AirPuffConfig | ImagingConfig | SampleChamberConfig | ProbeConfig | EphysAssemblyConfig | CatheterConfig ] = Field( ..., title="Device configurations", description="Configurations are parameters controlling active devices during this stream", ) connections: List[Connection] = Field( default=[], title="Connections", description=( "Connections are links between devices that are specific to this acquisition (i.e." " not already defined in the Instrument)" ), ) @model_validator(mode="after") def check_modality_config_requirements(self): """Check that the required devices are present for the modalities""" for modality in self.modalities: if modality.abbreviation not in CONFIG_REQUIREMENTS.keys(): # No configuration requirements for this modality continue # pragma: no cover for group in CONFIG_REQUIREMENTS[modality.abbreviation]: if not any(isinstance(config, device_type) for config in self.configurations for device_type in group): # Get the types of all configurations group_types = [device_type.__name__ for device_type in group] config_types = [type(config).__name__ for config in self.configurations] raise ValueError( f"Missing one of required devices {group_types} for modality {modality.name} in {config_types}" ) return self @model_validator(mode="after") def check_connections(self): """Check that every device in a Connection is present in the active_devices list""" for connection in self.connections: # Check that both source and target devices are in active_devices if ( connection.source_device not in self.active_devices or connection.target_device not in self.active_devices ): missing_devices = [] if connection.source_device not in self.active_devices: missing_devices.append(connection.source_device) if connection.target_device not in self.active_devices: missing_devices.append(connection.target_device) raise ValueError( f"Missing devices in active_devices list for connection " f"from '{connection.source_device}' to '{connection.target_device}': {missing_devices}" ) return self @classmethod def overlapping(cls, stream1: "DataStream", stream2: "DataStream", overlap_s: int) -> bool: """Check if two DataStream objects have overlapping start and end times""" start_diff = abs((stream1.stream_start_time - stream2.stream_start_time).total_seconds()) end_diff = abs((stream1.stream_end_time - stream2.stream_end_time).total_seconds()) return start_diff <= overlap_s and end_diff <= overlap_s def __add__(self, other: "DataStream", overlap_s: int = 120) -> "DataStream": """Combine two DataStream objects""" if not DataStream.overlapping(self, other, overlap_s=overlap_s): raise ValueError("Cannot combine DataStreams with non-overlapping start and end times.") min_start_time = min(self.stream_start_time, other.stream_start_time) max_end_time = max(self.stream_end_time, other.stream_end_time) # Combine modalities modalities = self.modalities + other.modalities modalities = remove_duplicates(modalities) # Combine active devices active_devices = self.active_devices + other.active_devices len_orig_devices = len(active_devices) active_devices = remove_duplicates(active_devices) if len(active_devices) < len_orig_devices: logger.warning( "Duplicate active devices were removed. Only DAQ devices should be shared in overlapped " "DataStreams." ) # Combine configurations configurations = self.configurations + other.configurations # Combine connections connections = self.connections + other.connections # Combine notes notes = merge_notes(self.notes, other.notes) return DataStream( stream_start_time=min_start_time, stream_end_time=max_end_time, modalities=modalities, active_devices=active_devices, configurations=configurations, connections=connections, notes=notes, ) class ExternalDataStream(DataModel): """A simplified data stream for acquisitions where instrument metadata is unavailable.""" object_type: Literal["ExternalDataStream"] = "ExternalDataStream" stream_start_time: Annotated[ AwareDatetimeWithDefault, Field(..., title="Stream start time"), TimeValidation.BETWEEN, ] stream_end_time: Annotated[ AwareDatetimeWithDefault, Field(..., title="Stream stop time"), TimeValidation.BETWEEN, ] modalities: List[Modality.ONE_OF] = Field( ..., title="Modalities", description="Modalities that are acquired in this stream" ) notes: Optional[str] = Field(default=None, title="Notes") class StimulusEpoch(DataModel): """All stimuli being presented to the subject. starting and stopping at approximately the same time. Not all acquisitions have StimulusEpochs. """ stimulus_start_time: Annotated[AwareDatetimeWithDefault, TimeValidation.BETWEEN] = Field( ..., title="Stimulus start time", description="When a specific stimulus begins. This might be the same as the acquisition start time.", ) stimulus_end_time: Annotated[AwareDatetimeWithDefault, TimeValidation.BETWEEN] = Field( ..., title="Stimulus end time", description="When a specific stimulus ends. This might be the same as the acquisition end time.", ) stimulus_name: str = Field(..., title="Stimulus name") code: Optional[Code] = Field( default=None, title="Code or script", description=( "Custom code/script used to control the behavior/stimulus." " Use the Code.parameters field to store stimulus properties" ), ) stimulus_modalities: List[StimulusModality] = Field(..., title="Stimulus modalities") performance_metrics: Optional[PerformanceMetrics] = Field(default=None, title="Performance metrics") notes: Optional[str] = Field(default=None, title="Notes") # Devices and configurations active_devices: List[str] = Field( default=[], title="Active devices", description="Device names must match devices in the Instrument", ) configurations: DiscriminatedList[ SpeakerConfig | LightEmittingDiodeConfig | LaserConfig | MousePlatformConfig | OlfactometerConfig ] = Field(default=[], title="Device configurations") # Training protocol training_protocol_name: Optional[str] = Field( default=None, title="Training protocol name", description=( "Name of the training protocol used during the acquisition, " "must match a protocol in the Procedures" ), ) curriculum_status: Optional[str] = Field( default=None, title="Curriculum status", description="Status within the training protocol curriculum", ) class Manipulation(ProtocolListMixin, DataModel): """Description of procedures performed during an acquisition.""" start_time: Annotated[AwareDatetimeWithDefault, TimeValidation.BETWEEN] = Field( ..., title="Manipulation start time", description="Must be between the acquisition start and end times" ) end_time: Annotated[AwareDatetimeWithDefault, TimeValidation.BETWEEN] = Field( ..., title="Manipulation end time", description="Must be between the acquisition start and end times" ) procedures: Optional[DiscriminatedList[Injection | BrainInjection | Reagent]] = Field( default=None, title="Procedures", description="Procedures performed during the manipulation" ) anaesthesia: Optional[Anaesthetic] = Field(default=None, title="Anaesthesia") notes: Optional[str] = Field(default=None, title="Notes") class Acquisition(ProtocolListMixin, DataCoreModel): """Description of data acquisition metadata including streams, stimuli, and experimental setup. The acquisition metadata is split into two parallel pieces: the DataStream and the StimulusEpoch. At any given moment in time the active DataStream(s) represents all modalities of data being acquired, while the StimulusEpoch represents all stimuli being presented.""" # Meta metadata _DESCRIBED_BY_URL = DataCoreModel._DESCRIBED_BY_BASE_URL.default + "aind_data_schema/core/acquisition.py" describedBy: str = Field(default=_DESCRIBED_BY_URL, json_schema_extra={"const": _DESCRIBED_BY_URL}) schema_version: SkipValidation[Literal["2.5.2"]] = Field(default="2.5.2") # ID subject_id: str = Field(default=..., title="Subject ID", description="Unique identifier for the subject") specimen_id: Optional[Union[str, List[str]]] = Field( default=None, title="Specimen ID", description="Required for in vitro modalities. Standard format is {subject_id} with a _### suffix, as needed", ) # Acquisition metadata acquisition_start_time: AwareDatetimeWithDefault = Field( ..., title="Acquisition start time", description="During validation, timezone information will be moved into the acquisition_start_tz field.", ) acquisition_start_tz: Optional[Union[int, TimeZoneName]] = Field( default=None, title="Acquisition start timezone", description=( "Automatically populated by a validator based on acquisition_start_time. " "Will be a TimeZoneName (IANA name) when the datetime uses a ZoneInfo timezone, " "or an integer UTC offset in hours for fixed-offset timezones. " "Use ZoneInfo (from the zoneinfo standard library) to preserve the named timezone." ), ) acquisition_end_time: AwareDatetimeWithDefault = Field(..., title="Acquisition end time") @field_validator("acquisition_start_tz", mode="before") @classmethod def coerce_fixed_offset_tz_string(cls, v): """Convert legacy fixed-offset strings like '-07:00' or '+05:30' to integer minutes.""" if isinstance(v, str): m = re.fullmatch(r"([+-]?)(\d{2}):(\d{2})", v) if m: sign = -1 if m.group(1) == "-" else 1 return sign * (int(m.group(2)) * 60 + int(m.group(3))) // 60 return v experimenters: List[str] = Field( default=[], title="experimenter(s)", ) ethics_review_id: Optional[List[str]] = Field(default=None, title="Ethics review ID") instrument_id: Optional[str] = Field( default=None, title="Instrument ID", description="Should match the Instrument.instrument_id. Required when instrument metadata is available.", ) acquisition_type: str = Field( ..., title="Acquisition type", description=( "Descriptive string detailing the type of acquisition, " "should be consistent across similar acquisitions for the same experiment." ), ) notes: Optional[str] = Field(default=None, title="Notes") # Coordinate system coordinate_system: Optional[CoordinateSystem] = Field( default=None, title="Coordinate system", description=( "Origin and axis definitions for determining the configured position of devices during acquisition." " Required when coordinates are provided within the Acquisition" ), ) # note: exact field name is used by a validator # Instrument metadata calibrations: List[CALIBRATIONS] = Field( default=[], title="Calibrations", description="List of calibration measurements taken prior to acquisition.", ) maintenance: List[Maintenance] = Field( default=[], title="Maintenance", description="List of maintenance on instrument prior to acquisition." ) # Acquisition data data_streams: DiscriminatedList[DataStream | ExternalDataStream] = Field( ..., title="Data streams", description=( "A data stream is a collection of devices that are acquiring data simultaneously. Each acquisition can " "include multiple streams. Streams should be split when configurations are changed. " "Use ExternalDataStream for acquisitions where instrument metadata is unavailable." ), ) stimulus_epochs: List[StimulusEpoch] = Field( default=[], title="Stimulus", description=( "A stimulus epoch captures all stimuli being presented during an acquisition." " Epochs should be split when the purpose of the stimulus changes." ), ) manipulations: List[Manipulation] = Field( default=[], title="Manipulations", description="Procedures performed during the acquisition." ) subject_details: Optional[AcquisitionSubjectDetails] = Field( default=None, title="Subject details", description="Required for in vivo acquisitions." ) @property def acquisition_start_time_local(self) -> datetime: """Return acquisition_start_time converted to the timezone stored in acquisition_start_tz. If acquisition_start_tz is a TimeZoneName (IANA name), uses ZoneInfo to construct the timezone. If acquisition_start_tz is an int (UTC offset in hours), uses a fixed-offset timezone. If acquisition_start_tz is None, returns acquisition_start_time as-is. """ if self.acquisition_start_tz is None: return self.acquisition_start_time if isinstance(self.acquisition_start_tz, int): tz = timezone(timedelta(hours=self.acquisition_start_tz)) else: tz = ZoneInfo(str(self.acquisition_start_tz)) return self.acquisition_start_time.astimezone(tz) @model_validator(mode="after") def extract_timezone(self): """Extract timezone information from acquisition_start_time and set acquisition_start_tz""" if self.acquisition_start_tz is None and hasattr(self, "acquisition_start_time"): self.acquisition_start_tz = extract_timezone_from_datetime(self.acquisition_start_time) return self @model_validator(mode="after") def check_subject_specimen_id(self): """Check that the subject and specimen IDs match""" if self.specimen_id and self.subject_id: ids = self.specimen_id if isinstance(self.specimen_id, list) else [self.specimen_id] for sid in ids: if not subject_specimen_id_compatibility(self.subject_id, sid): raise ValueError(f"Expected {self.subject_id} to appear in {sid}") return self @model_validator(mode="after") def instrument_id_required_for_data_streams(self): """Require instrument_id when any standard DataStream is present""" if not hasattr(self, "data_streams"): return self if any(isinstance(stream, DataStream) for stream in self.data_streams): if not self.instrument_id: raise ValueError("instrument_id is required when data_streams contains a DataStream") return self @model_validator(mode="after") def specimen_required(self): """Check if specimen ID is required for in vitro imaging modalities""" if not hasattr(self, "data_streams"): # bypass for testing return self for stream in self.data_streams: if any([modality.abbreviation in SPECIMEN_MODALITIES for modality in stream.modalities]): if not self.specimen_id: raise ValueError(f"Specimen ID is required for modalities {stream.modalities}") return self @classmethod def _merge_data_streams(cls, streams: List[DataStream], overlap_s: int = 120) -> List[DataStream]: """Merge two lists of data streams""" groups = [] visited = set() for i in range(len(streams)): if i in visited: continue group = [streams[i]] visited.add(i) for j in range(i + 1, len(streams)): if j not in visited and DataStream.overlapping(streams[i], streams[j], overlap_s=overlap_s): group.append(streams[j]) visited.add(j) groups.append(group) # Construct the final set of streams, including merged streams where applicable merged_streams = [] for group in groups: if len(group) == 1: merged_streams.append(group[0]) else: merged_stream = group[0] for stream in group[1:]: merged_stream = merged_stream + stream merged_streams.append(merged_stream) return merged_streams def __add__(self, other: "Acquisition") -> "Acquisition": """Combine two Acquisition objects""" # Check for schema version incompability if self.schema_version != other.schema_version: raise ValueError( "Cannot combine Acquisition objects with different schema " + f"versions: {self.schema_version} and {other.schema_version}" ) # Figure out what coordinate system to use coordinate_system = merge_coordinate_systems(self.coordinate_system, other.coordinate_system) # Check for incompatible key fields subj_check = self.subject_id != other.subject_id spec_check = self.specimen_id != other.specimen_id exp_type_check = self.acquisition_type != other.acquisition_type if any([subj_check, spec_check, exp_type_check]): raise ValueError( "Cannot combine Acquisition objects that differ in key fields:\n" f"subject_id: {self.subject_id}/{other.subject_id}\n" f"specimen_id: {self.specimen_id}/{other.specimen_id}\n" f"acquisition_type: {self.acquisition_type}/{other.acquisition_type}" ) # Combine instrument_id instrument_id = merge_str_alphabetical(self.instrument_id, other.instrument_id) details_check = self.subject_details and other.subject_details if details_check: raise ValueError( "SubjectDetails cannot be combined in Acquisition. Only a single set of details is allowed." ) # Combine experimenters = self.experimenters + other.experimenters protocol_id = merge_optional_list(self.protocol_id, other.protocol_id) ethics_review_id = merge_optional_list(self.ethics_review_id, other.ethics_review_id) calibrations = self.calibrations + other.calibrations maintenance = self.maintenance + other.maintenance all_streams = self.data_streams + other.data_streams external_streams = [s for s in all_streams if isinstance(s, ExternalDataStream)] regular_streams = [s for s in all_streams if isinstance(s, DataStream)] data_streams = Acquisition._merge_data_streams(regular_streams) + external_streams stimulus_epochs = self.stimulus_epochs + other.stimulus_epochs # Remove duplicates experimenters = remove_duplicates(experimenters) if ethics_review_id: ethics_review_id = remove_duplicates(ethics_review_id) # Combine notes notes = merge_notes(self.notes, other.notes) # Handle start and end time start_time = min(self.acquisition_start_time, other.acquisition_start_time) end_time = max(self.acquisition_end_time, other.acquisition_end_time) return Acquisition( subject_id=self.subject_id, specimen_id=self.specimen_id, experimenters=experimenters, protocol_id=protocol_id, ethics_review_id=ethics_review_id, instrument_id=instrument_id, calibrations=calibrations, coordinate_system=coordinate_system, maintenance=maintenance, acquisition_start_time=start_time, acquisition_end_time=end_time, acquisition_type=self.acquisition_type, notes=notes, data_streams=data_streams, stimulus_epochs=stimulus_epochs, subject_details=self.subject_details if self.subject_details else other.subject_details, )