From 719cc0d3dbaccee52549b2c1ed4f5c5b9bdfa022 Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Wed, 9 Sep 2026 20:17:57 -0700 Subject: [PATCH 1/2] adds get_metrics to composite trial --- .../trial_generators/composite_trial_generator.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) 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 29b74b91..57a2af36 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, TrialOutcome, TrialMetrics 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() From 02080047ebba0f1fb89ec2611750e0d6a56774b1 Mon Sep 17 00:00:00 2001 From: Micah Woodard Date: Thu, 10 Sep 2026 08:53:34 -0700 Subject: [PATCH 2/2] lints --- .../task_logic/trial_generators/composite_trial_generator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 57a2af36..94cdd3bd 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, TrialMetrics +from ..trial_models import Trial, TrialMetrics, TrialOutcome from ._base import BaseTrialGeneratorSpecModel, ITrialGenerator _TSpec = TypeVar("_TSpec", bound=BaseTrialGeneratorSpecModel, covariant=True)