diff --git a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/composite_trial_generator.py b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/composite_trial_generator.py index 29b74b9..94cdd3b 100644 --- a/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/composite_trial_generator.py +++ b/src/aind_behavior_dynamic_foraging/task_logic/trial_generators/composite_trial_generator.py @@ -2,7 +2,7 @@ from pydantic import Field, SerializeAsAny -from ..trial_models import Trial, TrialOutcome +from ..trial_models import Trial, TrialMetrics, TrialOutcome from ._base import BaseTrialGeneratorSpecModel, ITrialGenerator _TSpec = TypeVar("_TSpec", bound=BaseTrialGeneratorSpecModel, covariant=True) @@ -70,3 +70,8 @@ def update(self, outcome: TrialOutcome | str) -> None: """ if self._current_index < len(self._generators): self._generators[self._current_index].update(outcome) + + def get_metrics(self) -> TrialMetrics: + """Return metrics at current state of the trial generator.""" + if self._current_index < len(self._generators): + return self._generators[self._current_index].get_metrics()