From e1ec0691e442535e6a4a8ce97bfd0c24617eb93e Mon Sep 17 00:00:00 2001 From: ttaleh Date: Mon, 5 Oct 2026 16:03:44 +0300 Subject: [PATCH 1/3] Fix batching with random task durations A case enters a batched task's waiting list when its previous task is computed, which happens ahead of time, so the list holds cases that reach the task only later and is not ordered by time. - After a batch fires, remove from the waiting list exactly the cases the batch takes, not the first N in insertion order. Before, a case could run twice (crash) while another one was lost. - The firing rule and the batch size only look at cases that have reached the batched task by now, in the order they reached it, so they agree on which cases are waiting. - Use total_seconds() instead of .seconds for waiting times; .seconds wraps negative differences and drops whole days. Cases that haven't reached the batched task are no longer swept into an earlier batch, so the last cases of a run are now fired by the end of the run, which needs two more fixes: - At the end of the run, waiting cases whose last gap is below the low boundary get a firing time instead of being lost, and the waiting list is sorted by the time cases reached the task. - A single waiting case under large_wt waits up to the high boundary instead of firing on the low one. --- prosimos/batch_processing.py | 29 +++++-- prosimos/control_flow_manager.py | 78 +++++++++--------- .../test_batching_random_durations.py | 79 +++++++++++++++++++ 3 files changed, 144 insertions(+), 42 deletions(-) create mode 100644 testing_scripts/test_batching_random_durations.py diff --git a/prosimos/batch_processing.py b/prosimos/batch_processing.py index 4de0143..ceec0bd 100644 --- a/prosimos/batch_processing.py +++ b/prosimos/batch_processing.py @@ -123,7 +123,7 @@ def is_true(self, element): curr_enabled_datetime = element["curr_enabled_at"] op = _get_operator_symbols_ge(self.operator) - ready_wt_sec = (curr_enabled_datetime - last_enabled_datetime).seconds + ready_wt_sec = (curr_enabled_datetime - last_enabled_datetime).total_seconds() is_rule_true = op(ready_wt_sec, self.value2) if is_rule_true == False: @@ -466,7 +466,7 @@ def get_ready_wt(self, element): # happens when no new cases will arrive return en_time_index, _get_enabled_time_for_wt_rule(prev_item, operator.gt, high_boundary) - diff = (item - prev_item).seconds + diff = (item - prev_item).total_seconds() is_batch_enabled_low = diff < low_boundary if is_batch_enabled_low: @@ -508,14 +508,18 @@ def get_large_wt(self, element): prev_item = first_item list_len = len(enabled_dt_with_curr_enabled) for en_time_index, item in enumerate(enabled_dt_with_curr_enabled[1:], 1): - diff = (item - first_item).seconds + diff = (item - first_item).total_seconds() is_batch_enabled_low = diff > low_boundary if not is_batch_enabled_low: if en_time_index == list_len - 1: # last item of the evaluation # this will enable the batch in the future - result = en_time_index, _get_enabled_time_for_wt_rule(first_item, operator.ge, low_boundary) + if en_time_index == 1: + # a single waiting case waits for a second one up to the high boundary + result = en_time_index, _get_enabled_time_for_wt_rule(first_item, operator.gt, high_boundary) + else: + result = en_time_index, _get_enabled_time_for_wt_rule(first_item, operator.ge, low_boundary) break prev_item = item @@ -565,6 +569,21 @@ def _get_min_enabled_time_waiting_time(self, case_id_and_enabled_times, wt_res = self.get_ready_wt(draft_element) \ if rule_type == RULE_TYPE.READY_WT else self.get_large_wt(draft_element) + if wt_res == (0, None): + # no new case reaches the batched task anymore, so: + # - a single waiting case fires once it has waited past the high boundary; + # - several waiting cases fire once the low boundary has passed, counted from + # the last case's arrival for ready_wt (time since the last arrival), + # or from the first case's arrival for large_wt (time the oldest case has waited) + low_boundary, high_boundary = self.ready_wt_boundaries \ + if rule_type == RULE_TYPE.READY_WT else self.large_wt_boundaries + enabled_datetimes = draft_element["enabled_datetimes"] + if draft_element["size"] == 1: + wt_res = 1, _get_enabled_time_for_wt_rule(enabled_datetimes[0], operator.gt, high_boundary) + else: + reference = enabled_datetimes[-1] if rule_type == RULE_TYPE.READY_WT else enabled_datetimes[0] + wt_res = draft_element["size"], _get_enabled_time_for_wt_rule(reference, operator.ge, low_boundary) + if wt_res == None: return None else: @@ -913,7 +932,7 @@ def is_invalid_end(self, num_tasks, first_wt, ready_wt): Check whether items waiting for batch execution might be satisfied in the future (valid for further processing) or they are invalid (one part of the AND rule could not be satisfied in the future at all) :param num_tasks: number of tasks waiting for batch execution - :param first_wt: waiting time of the first item (current_point_in_time - first_item.enable_time).seconds + :param first_wt: waiting time of the first item (current_point_in_time - first_item.enable_time).total_seconds() :param ready_wt: waiting time of the last task in the batch queue :return: whether the rule is invalid :rtype: boolean diff --git a/prosimos/control_flow_manager.py b/prosimos/control_flow_manager.py index ea509f9..5826098 100644 --- a/prosimos/control_flow_manager.py +++ b/prosimos/control_flow_manager.py @@ -19,9 +19,9 @@ seconds_per_unit = {"s": 1, "m": 60, "h": 3600, "d": 86400, "w": 604800} class BatchInfoForExecution: - def __init__(self, all_case_ids, task_batch_info, curr_task_id, batch_spec, start_time_from_rule): - self.case_ids = all_case_ids[curr_task_id].copy() - self.task_batch_info = task_batch_info[curr_task_id] + def __init__(self, case_ids, task_batch_info, batch_spec, start_time_from_rule): + self.case_ids = case_ids + self.task_batch_info = task_batch_info self.batch_spec = batch_spec self.start_time_from_rule = start_time_from_rule self.batch_id = str(uuid.uuid4()) @@ -400,18 +400,19 @@ def is_batched_task_enabled(self, task_id: str, enabled_at: CustomDatetimeAndSec firing_rules: List[AndFiringRule] = task_batch_info.firing_rules - size_count = self.batch_count[task_id] if self.batch_count.get(task_id, None) != None else 0 + reached = self._reached_batch(task_id, enabled_at) + size_count = len(reached) if not firing_rules.is_batch_size_enough_for_exec(size_count): #size_count < 2: # not enough items for batch execution return False, None, None - waiting_time = [ (enabled_at.datetime - v.datetime).total_seconds() for (_, v) in self.batch_waiting_processes[task_id].items() ] + waiting_time = [ (enabled_at.datetime - v.datetime).total_seconds() for (_, v) in reached ] spec = { "size": size_count, "waiting_times": waiting_time, - "enabled_datetimes": [ v.datetime for (_, v) in self.batch_waiting_processes[task_id].items() ], + "enabled_datetimes": [ v.datetime for (_, v) in reached ], "curr_enabled_at": enabled_at.datetime, "is_triggered_by_batch": True, # specify where from we checking the rule. If not triggered by batch - then we move to midnight time "is_only_one_batch_return": False @@ -425,6 +426,16 @@ def is_batched_task_enabled(self, task_id: str, enabled_at: CustomDatetimeAndSec else: return self.get_batch_size_no_rules_defined(spec, firing_rules, task_batch_info) + def _reached_batch(self, task_id: str, enabled_at: CustomDatetimeAndSeconds): + """ + Cases waiting for the batched task that have reached it by enabled_at, in the order they reached it. + A case is added to the waiting list when its previous task is computed, which happens ahead of time, + so the list also holds cases that reach the batched task only later, and it is not ordered by time. + """ + reached = [ (case_id, reached_at) for (case_id, reached_at) in self.batch_waiting_processes[task_id].items() + if reached_at.datetime <= enabled_at.datetime ] + return sorted(reached, key=lambda item: item[1].datetime) + def get_batch_size_no_rules_defined(self, spec, firing_rules, task_batch_info): """ In case no rules were discovered, size distribution is being used to define how batches will be formed. @@ -471,7 +482,8 @@ def get_batch_size_no_rules_defined(self, spec, firing_rules, task_batch_info): def get_start_time(self, task_id, last_task_enabled_time) -> tuple([int, int, CustomDatetimeAndSeconds]): task_batch_info = self.batch_info.get(task_id, None) firing_rules: List[AndFiringRule] = task_batch_info.firing_rules - enabled_times = list(self.batch_waiting_processes[task_id].items()) + # the waiting list is not ordered by the time cases reached the batched task + enabled_times = sorted(self.batch_waiting_processes[task_id].items(), key=lambda item: item[1].datetime) batch_enabled_time = firing_rules.get_enabled_time( enabled_times, last_task_enabled_time, @@ -939,35 +951,23 @@ def move_batch_if_enabled(self, task_id: str, case_id, enabled_time: CustomDatet is_enabled, batch_spec, start_time_from_rule = self.is_batched_task_enabled(task_id, enabled_time) if is_enabled: - batch_info = BatchInfoForExecution( - self.batch_waiting_processes, - self.batch_info, - task_id, - batch_spec, - start_time_from_rule) + batch_info = self.create_batch_info_and_clear_from_queue( + task_id, batch_spec, start_time_from_rule, self._reached_batch(task_id, enabled_time) + ) enabled_tasks.append((EnabledTask(task_id, batch_info))) - self._clear_batch(task_id, batch_spec) - def _clear_batch(self, next_e, batch_spec): + def _clear_batch(self, next_e, case_ids): """ When we passed on the information about the batch for the execution, clear that data from here to avoid multiple execution """ - for batch_size in batch_spec: - if batch_size != None: - # remove first batch_size-element since they are being executed - # the rest stays in the queue for being enabled for batch execution - curr_index = 0 - for item_key in list(self.batch_waiting_processes[next_e].keys()): - del self.batch_waiting_processes[next_e][item_key] - curr_index = curr_index + 1 + # remove exactly the cases being executed + # the rest stays in the queue for being enabled for batch execution + for case_id in case_ids: + del self.batch_waiting_processes[next_e][case_id] - if curr_index == batch_size: - # all waiting processes regarding the selected batch was removed - break - - self.batch_count[next_e] = self.batch_count[next_e] - batch_size + self.batch_count[next_e] = self.batch_count[next_e] - len(case_ids) def increase_task_count(self, task_id, case_id, enabled_time): @@ -998,7 +998,7 @@ def is_any_batch_enabled(self, started_datetime: CustomDatetimeAndSeconds): is_enabled, batch_spec, start_time_from_rule = self.is_batched_task_enabled(task_id, started_datetime) if is_enabled: enabled_task_batch[task_id] = self.create_batch_info_and_clear_from_queue( - task_id, batch_spec, start_time_from_rule + task_id, batch_spec, start_time_from_rule, self._reached_batch(task_id, started_datetime) ) return enabled_task_batch @@ -1042,30 +1042,34 @@ def get_invalid_batches_if_any(self, current_point_of_time): last_task_in_batch_start = waiting_tasks[last_added_key].datetime start_time_from_rule = max(current_point_of_time.datetime, last_task_in_batch_start) enabled_task_batch[task_id] = self.create_batch_info_and_clear_from_queue( - task_id, batch_spec, start_time_from_rule + task_id, batch_spec, start_time_from_rule, list(waiting_tasks.items()) ) continue return enabled_task_batch - def create_batch_info_and_clear_from_queue(self, task_id: str, batch_spec, start_time_from_rule): + def create_batch_info_and_clear_from_queue(self, task_id: str, batch_spec, start_time_from_rule, candidates): + """ + :param candidates: (case_id, enabled_time) pairs the firing rule looked at, in the same order. + The batch takes the first of them, as many as batch_spec counts in total + """ + case_ids = dict(candidates[:sum(batch_spec)]) batch_info = BatchInfoForExecution( - self.batch_waiting_processes, - self.batch_info, - task_id, + case_ids, + self.batch_info[task_id], batch_spec, start_time_from_rule) - self._clear_batch(task_id, batch_spec) + self._clear_batch(task_id, case_ids) return batch_info def is_or_rule_invalid(self, waiting_tasks, task_id: str, num_tasks_wait_batch: int, current_point_of_time: CustomDatetimeAndSeconds): all_keys = list(waiting_tasks.keys()) first_key = all_keys[0] - first_wt = (current_point_of_time.datetime - waiting_tasks[first_key].datetime).seconds + first_wt = (current_point_of_time.datetime - waiting_tasks[first_key].datetime).total_seconds() last_key = all_keys[-1] - last_wt = (current_point_of_time.datetime - waiting_tasks[last_key].datetime).seconds + last_wt = (current_point_of_time.datetime - waiting_tasks[last_key].datetime).total_seconds() return self.batch_info[task_id].firing_rules.is_invalid_end(num_tasks_wait_batch, first_wt, last_wt) diff --git a/testing_scripts/test_batching_random_durations.py b/testing_scripts/test_batching_random_durations.py new file mode 100644 index 0000000..41257ae --- /dev/null +++ b/testing_scripts/test_batching_random_durations.py @@ -0,0 +1,79 @@ +import json +import random + +import numpy as np +import pandas as pd +import pytest + +from prosimos.simulation_engine import run_simulation +from testing_scripts.test_batching import assets_path + +MODEL_FILENAME = "batch-example-end-task.bpmn" +JSON_FILENAME = "batch-example-with-batch.json" +BATCHED_TASK = "D" +TOTAL_CASES = 40 + +# other batching tests overwrite these two sections of the example file in place, +# so they are set back to the committed values here +COMMITTED_ARRIVAL_DISTRIBUTION = { + "distribution_name": "expon", + "distribution_params": [{"value": 7200.0}, {"value": 0.0}, {"value": 10000.0}], +} +COMMITTED_BATCH_PROCESSING = [ + { + "task_id": "sid-503A048D-6344-446A-8D67-172B164CF8FA", + "type": "Parallel", + "batch_frequency": 1.0, + "size_distrib": [{"key": "1", "value": 0}, {"key": "2", "value": 1}], + "duration_distrib": [{"key": "3", "value": 0.8}], + "firing_rules": [ + [ + {"attribute": "ready_wt", "comparison": ">", "value": 7200}, + {"attribute": "ready_wt", "comparison": "<", "value": 10800}, + ] + ], + } +] + + +def _committed_example(assets_path): + with open(assets_path / JSON_FILENAME) as file: + settings = json.load(file) + settings["arrival_time_distribution"] = COMMITTED_ARRIVAL_DISTRIBUTION + settings["batch_processing"] = json.loads(json.dumps(COMMITTED_BATCH_PROCESSING)) + return settings + + +def _run(assets_path, tmp_path, settings): + json_path = tmp_path / "settings.json" + with open(json_path, "w") as file: + json.dump(settings, file) + log_path = tmp_path / "log.csv" + + random.seed(1) + np.random.seed(1) + run_simulation(assets_path / MODEL_FILENAME, json_path, TOTAL_CASES, None, log_path, + "2024-01-01T09:00:00+00:00") + + log = pd.read_csv(log_path) + return log[log["activity"] == BATCHED_TASK].groupby("case_id").size() + + +@pytest.mark.parametrize("batch_type", ["Parallel", "Sequential"]) +def test_random_durations_run_every_case_through_the_batch_once(assets_path, tmp_path, batch_type): + # ====== ARRANGE ====== + # random task durations make cases reach the batched task out of the order they were added to its queue, + # and the queue holds cases that reach it only later on + settings = _committed_example(assets_path) + for task in settings["task_resource_distribution"]: + for resource in task["resources"]: + resource["distribution_name"] = "expon" + resource["distribution_params"] = [{"value": 1800}, {"value": 0}, {"value": 18000}] + settings["batch_processing"][0]["type"] = batch_type + + # ====== ACT ====== + runs_per_case = _run(assets_path, tmp_path, settings) + + # ====== ASSERT ====== + assert len(runs_per_case) == TOTAL_CASES + assert (runs_per_case == 1).all() From c83d59d9197f6a7e2325243eea5aac33abe540c4 Mon Sep 17 00:00:00 2001 From: ttaleh Date: Mon, 5 Oct 2026 16:04:05 +0300 Subject: [PATCH 2/3] Stop the false "batch size ... 0" warning for a single waiting case With a ready_wt rule, the batch size has always made a single waiting case wait for a second one up to the high boundary, but the firing rule said it should fire once the low boundary passed. The two disagreed, nothing fired, and the "batch size ... 0" warning was printed, also with fixed task durations. The low-boundary part of the rule now stays false for a single waiting case, so it agrees with the batch size. --- prosimos/batch_processing.py | 5 +++++ .../test_batching_random_durations.py | 21 ++++++++++++++++++- 2 files changed, 25 insertions(+), 1 deletion(-) diff --git a/prosimos/batch_processing.py b/prosimos/batch_processing.py index ceec0bd..cb5994e 100644 --- a/prosimos/batch_processing.py +++ b/prosimos/batch_processing.py @@ -119,6 +119,11 @@ def is_true(self, element): return is_rule_true elif self.variable1 == "ready_wt": + if queue_size == 1 and self.operator in [">", ">="]: + # a single waiting case does not fire on the low boundary: + # it waits for a second case up to the high boundary, as get_ready_wt counts it + return False + last_enabled_datetime = element["enabled_datetimes"][-1] curr_enabled_datetime = element["curr_enabled_at"] op = _get_operator_symbols_ge(self.operator) diff --git a/testing_scripts/test_batching_random_durations.py b/testing_scripts/test_batching_random_durations.py index 41257ae..17b0b46 100644 --- a/testing_scripts/test_batching_random_durations.py +++ b/testing_scripts/test_batching_random_durations.py @@ -12,6 +12,7 @@ JSON_FILENAME = "batch-example-with-batch.json" BATCHED_TASK = "D" TOTAL_CASES = 40 +BATCH_SIZE_ZERO_WARNING = "batch size for the execution returned to be 0" # other batching tests overwrite these two sections of the example file in place, # so they are set back to the committed values here @@ -60,7 +61,7 @@ def _run(assets_path, tmp_path, settings): @pytest.mark.parametrize("batch_type", ["Parallel", "Sequential"]) -def test_random_durations_run_every_case_through_the_batch_once(assets_path, tmp_path, batch_type): +def test_random_durations_run_every_case_through_the_batch_once(assets_path, tmp_path, capsys, batch_type): # ====== ARRANGE ====== # random task durations make cases reach the batched task out of the order they were added to its queue, # and the queue holds cases that reach it only later on @@ -77,3 +78,21 @@ def test_random_durations_run_every_case_through_the_batch_once(assets_path, tmp # ====== ASSERT ====== assert len(runs_per_case) == TOTAL_CASES assert (runs_per_case == 1).all() + assert BATCH_SIZE_ZERO_WARNING not in capsys.readouterr().out + + +@pytest.mark.parametrize("batch_type", ["Parallel", "Sequential"]) +def test_fixed_durations_print_no_batch_size_zero_warning(assets_path, tmp_path, capsys, batch_type): + # ====== ARRANGE ====== + # the example as committed: a single waiting case between the ready_wt boundaries + # used to make the firing rule and the batch size disagree + settings = _committed_example(assets_path) + settings["batch_processing"][0]["type"] = batch_type + + # ====== ACT ====== + runs_per_case = _run(assets_path, tmp_path, settings) + + # ====== ASSERT ====== + assert len(runs_per_case) == TOTAL_CASES + assert (runs_per_case == 1).all() + assert BATCH_SIZE_ZERO_WARNING not in capsys.readouterr().out From 621b47115a1a26b2a95f475012815cb32d8d8505 Mon Sep 17 00:00:00 2001 From: ttaleh Date: Mon, 5 Oct 2026 16:31:49 +0300 Subject: [PATCH 3/3] Round batch waiting times down to whole seconds The firing rule boundaries are whole seconds (e.g. "> 7200" becomes a low boundary of 7201). total_seconds() gives fractions, so a wait of 7200.6 s passed "> 7200" in the firing rule while the batch size saw it below 7201 and returned 0, printing the "batch size ... 0" warning. timedelta.seconds rounded down but dropped whole days. whole_seconds() rounds down like .seconds and keeps the days; it is used for the waiting times compared with the boundaries. The batching tests now also put the global random generators' states back after each test, so later tests draw the same values as without them. --- prosimos/batch_processing.py | 15 ++-- prosimos/control_flow_manager.py | 6 +- .../test_batching_random_durations.py | 68 +++++++++++++++++-- 3 files changed, 76 insertions(+), 13 deletions(-) diff --git a/prosimos/batch_processing.py b/prosimos/batch_processing.py index cb5994e..60af00d 100644 --- a/prosimos/batch_processing.py +++ b/prosimos/batch_processing.py @@ -34,6 +34,13 @@ def _get_operator_symbols_lt(operator_str: str): def _is_greater(op: operator): return op in [operator.ge, operator.gt] +def whole_seconds(difference: timedelta) -> int: + """ + Seconds in a time difference, rounded down to a whole second, days included. + The rule boundaries are whole seconds; timedelta.seconds rounds down too but drops the days. + """ + return difference // timedelta(seconds=1) + class BATCH_TYPE(Enum): SEQUENTIAL = 'Sequential' # one after another CONCURRENT = 'Concurrent' # tasks are in progress simultaneously @@ -128,7 +135,7 @@ def is_true(self, element): curr_enabled_datetime = element["curr_enabled_at"] op = _get_operator_symbols_ge(self.operator) - ready_wt_sec = (curr_enabled_datetime - last_enabled_datetime).total_seconds() + ready_wt_sec = whole_seconds(curr_enabled_datetime - last_enabled_datetime) is_rule_true = op(ready_wt_sec, self.value2) if is_rule_true == False: @@ -471,7 +478,7 @@ def get_ready_wt(self, element): # happens when no new cases will arrive return en_time_index, _get_enabled_time_for_wt_rule(prev_item, operator.gt, high_boundary) - diff = (item - prev_item).total_seconds() + diff = whole_seconds(item - prev_item) is_batch_enabled_low = diff < low_boundary if is_batch_enabled_low: @@ -513,7 +520,7 @@ def get_large_wt(self, element): prev_item = first_item list_len = len(enabled_dt_with_curr_enabled) for en_time_index, item in enumerate(enabled_dt_with_curr_enabled[1:], 1): - diff = (item - first_item).total_seconds() + diff = whole_seconds(item - first_item) is_batch_enabled_low = diff > low_boundary if not is_batch_enabled_low: @@ -937,7 +944,7 @@ def is_invalid_end(self, num_tasks, first_wt, ready_wt): Check whether items waiting for batch execution might be satisfied in the future (valid for further processing) or they are invalid (one part of the AND rule could not be satisfied in the future at all) :param num_tasks: number of tasks waiting for batch execution - :param first_wt: waiting time of the first item (current_point_in_time - first_item.enable_time).total_seconds() + :param first_wt: waiting time of the first item, whole_seconds(current_point_in_time - first_item.enable_time) :param ready_wt: waiting time of the last task in the batch queue :return: whether the rule is invalid :rtype: boolean diff --git a/prosimos/control_flow_manager.py b/prosimos/control_flow_manager.py index 5826098..ca305c3 100644 --- a/prosimos/control_flow_manager.py +++ b/prosimos/control_flow_manager.py @@ -10,7 +10,7 @@ from pix_framework.statistics.distribution import DurationDistribution from prosimos.batch_processing import (BATCH_TYPE, AndFiringRule, - BatchConfigPerTask) + BatchConfigPerTask, whole_seconds) from prosimos.exceptions import InvalidBpmnModelException from prosimos.weekday_helper import CustomDatetimeAndSeconds from prosimos.simulation_execution_stats import SimulationExecutionStats @@ -1067,9 +1067,9 @@ def is_or_rule_invalid(self, waiting_tasks, task_id: str, num_tasks_wait_batch: all_keys = list(waiting_tasks.keys()) first_key = all_keys[0] - first_wt = (current_point_of_time.datetime - waiting_tasks[first_key].datetime).total_seconds() + first_wt = whole_seconds(current_point_of_time.datetime - waiting_tasks[first_key].datetime) last_key = all_keys[-1] - last_wt = (current_point_of_time.datetime - waiting_tasks[last_key].datetime).total_seconds() + last_wt = whole_seconds(current_point_of_time.datetime - waiting_tasks[last_key].datetime) return self.batch_info[task_id].firing_rules.is_invalid_end(num_tasks_wait_batch, first_wt, last_wt) diff --git a/testing_scripts/test_batching_random_durations.py b/testing_scripts/test_batching_random_durations.py index 17b0b46..e3c0ce8 100644 --- a/testing_scripts/test_batching_random_durations.py +++ b/testing_scripts/test_batching_random_durations.py @@ -1,10 +1,12 @@ import json import random +from datetime import datetime, timedelta, timezone import numpy as np import pandas as pd import pytest +from prosimos.batch_processing import AndFiringRule, FiringSubRule from prosimos.simulation_engine import run_simulation from testing_scripts.test_batching import assets_path @@ -37,6 +39,16 @@ ] +@pytest.fixture(autouse=True) +def _keep_random_state(): + # these tests seed the global random generators; their states are put back afterwards, + # so later tests draw the same values as without these tests + python_state, numpy_state = random.getstate(), np.random.get_state() + yield + random.setstate(python_state) + np.random.set_state(numpy_state) + + def _committed_example(assets_path): with open(assets_path / JSON_FILENAME) as file: settings = json.load(file) @@ -45,14 +57,14 @@ def _committed_example(assets_path): return settings -def _run(assets_path, tmp_path, settings): +def _run(assets_path, tmp_path, settings, seed=1): json_path = tmp_path / "settings.json" with open(json_path, "w") as file: json.dump(settings, file) log_path = tmp_path / "log.csv" - random.seed(1) - np.random.seed(1) + random.seed(seed) + np.random.seed(seed) run_simulation(assets_path / MODEL_FILENAME, json_path, TOTAL_CASES, None, log_path, "2024-01-01T09:00:00+00:00") @@ -81,18 +93,62 @@ def test_random_durations_run_every_case_through_the_batch_once(assets_path, tmp assert BATCH_SIZE_ZERO_WARNING not in capsys.readouterr().out +@pytest.mark.parametrize("seed", [1, 8]) @pytest.mark.parametrize("batch_type", ["Parallel", "Sequential"]) -def test_fixed_durations_print_no_batch_size_zero_warning(assets_path, tmp_path, capsys, batch_type): +def test_fixed_durations_print_no_batch_size_zero_warning(assets_path, tmp_path, capsys, batch_type, seed): # ====== ARRANGE ====== - # the example as committed: a single waiting case between the ready_wt boundaries + # the example as committed: a single waiting case between the ready_wt boundaries (seed 1), + # or a wait a fraction of a second past a boundary (seed 8), # used to make the firing rule and the batch size disagree settings = _committed_example(assets_path) settings["batch_processing"][0]["type"] = batch_type # ====== ACT ====== - runs_per_case = _run(assets_path, tmp_path, settings) + runs_per_case = _run(assets_path, tmp_path, settings, seed) # ====== ASSERT ====== assert len(runs_per_case) == TOTAL_CASES assert (runs_per_case == 1).all() assert BATCH_SIZE_ZERO_WARNING not in capsys.readouterr().out + + +@pytest.mark.parametrize( + "wait_sec, expected_batch_size", + [ + (7200.6, 0), # rounded down to 7200: not past "> 7200" yet + (7201.4, 2), # rounded down to 7201: past "> 7200", both cases fire + ], +) +def test_fraction_of_a_second_past_a_boundary_gives_the_same_answer_in_rule_and_batch_size( + capsys, wait_sec, expected_batch_size +): + # ====== ARRANGE ====== + # the boundaries are whole seconds, and a wait a fraction of a second past one + # used to make the firing rule say "fire" while the batch size said 0 + rule = AndFiringRule([FiringSubRule("ready_wt", ">", 7200), FiringSubRule("ready_wt", "<", 10800)]) + rule.init_boundaries() + first = datetime(2024, 1, 2, 9, 0, 0, tzinfo=timezone.utc) + last = first + timedelta(minutes=10) + now = last + timedelta(seconds=wait_sec) + + def element(): + return { + "size": 2, + "waiting_times": [(now - first).total_seconds(), (now - last).total_seconds()], + "enabled_datetimes": [first, last], + "curr_enabled_at": now, + "is_triggered_by_batch": True, + "is_only_one_batch_return": False, + } + + # ====== ACT ====== + rule_says_fire = all(subrule.is_true(element()) for subrule in rule.rules) + batch_size, _ = rule.get_firing_batch_size(2, element()) + is_true, batch_spec, _ = rule.is_true(element()) + + # ====== ASSERT ====== + assert rule_says_fire == (batch_size > 0) + assert batch_size == expected_batch_size + assert is_true == (expected_batch_size > 0) + assert batch_spec == ([expected_batch_size] if expected_batch_size else None) + assert BATCH_SIZE_ZERO_WARNING not in capsys.readouterr().out