feat: implement sub-block handling for shifts and update constraints in RotaBuilder

This commit is contained in:
Ross
2026-07-10 08:08:51 +01:00
parent d1fbbe2dc9
commit 5ff78faa07
3 changed files with 191 additions and 97 deletions
+128 -92
View File
@@ -417,7 +417,7 @@ class SingleShift(BaseModel):
rota_on_nwds: bool = False
assign_as_block: bool = False
force_as_block: bool = False
force_as_block_unless_nwd: bool = False
force_as_block_unless_nwd: bool | List[List[str]] = False
hard_constrain_shift: bool = True
bank_holidays_only: bool = False
#constraint: list[ShiftConstraint] = []
@@ -1050,20 +1050,28 @@ class RotaBuilder(object):
initialize=0,
)
blocks_worker_keys = []
for worker, week, shift_name in self.worker_week_shifts_to_assign_as_blocks():
shift = self.get_shift_by_name(shift_name)
sub_blocks = self.get_shift_sub_blocks(shift)
for sb_idx in range(len(sub_blocks)):
blocks_worker_keys.append((worker.id, week, shift_name, sb_idx))
self.model.blocks_worker_shift_assigned = Var(
(
(worker.id, week, shift)
for worker, week, shift in self.worker_week_shifts_to_assign_as_blocks()
),
blocks_worker_keys,
within=Binary,
initialize=0,
)
blocks_assigned_keys = []
for week, shift_name in self.week_shifts_to_assign_as_blocks():
shift = self.get_shift_by_name(shift_name)
sub_blocks = self.get_shift_sub_blocks(shift)
for sb_idx in range(len(sub_blocks)):
blocks_assigned_keys.append((week, shift_name, sb_idx))
self.model.blocks_assigned = Var(
(
(week, shift)
for week, shift in self.week_shifts_to_assign_as_blocks()
),
blocks_assigned_keys,
within=NonNegativeIntegers,
initialize=0,
)
@@ -1900,99 +1908,103 @@ class RotaBuilder(object):
for week, shift_name in self.week_shifts_to_assign_as_blocks():
shift = self.get_shift_by_name(shift_name)
sub_blocks = self.get_shift_sub_blocks(shift)
for worker in self.get_workers_for_shift(shift):
if (worker.id, week, shift.name) in self.model.blocks_worker_shift_assigned:
try:
self.model.constraints.add(
8
* self.model.blocks_worker_shift_assigned[
worker.id, week, shift_name
]
>= sum(
self.model.works[worker.id, week, day, shift_name]
for day in shift.days
for sb_idx, sb in enumerate(sub_blocks):
for worker in self.get_workers_for_shift(shift):
if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned:
try:
self.model.constraints.add(
8
* self.model.blocks_worker_shift_assigned[
worker.id, week, shift_name, sb_idx
]
>= sum(
self.model.works[worker.id, week, day, shift_name]
for day in sb
)
)
)
except KeyError:
pass
self.model.constraints.add(
self.model.blocks_assigned[week, shift_name]
== sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift_name
]
for worker in self.workers
if (worker.id, week, shift.name) in self.model.blocks_worker_shift_assigned
except KeyError:
pass
self.model.constraints.add(
self.model.blocks_assigned[week, shift_name, sb_idx]
== sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift_name, sb_idx
]
for worker in self.workers
if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned
)
)
)
workers_required = shift.get_worker_requirement_by_date(
self.get_week_start_date(week)
)
workers_required = shift.get_worker_requirement_by_date(
self.get_week_start_date(week)
)
if shift.force_as_block_unless_nwd:
workers = self.get_workers_for_shift(shift)
# Get workers who have a nwd on the shift
nwd_workers = []
full_workers = []
if shift.force_as_block_unless_nwd:
workers = self.get_workers_for_shift(shift)
# Get workers who have a nwd on the sub-block
nwd_workers = []
full_workers = []
for w in workers:
if w.non_working_day_list:
l = []
for (
nwd,
start_nwd_date,
end_nwd_date,
) in w.non_working_day_list:
if nwd in shift.days:
if start_nwd_date > self.get_week_start_date(week):
continue
if end_nwd_date < self.get_week_start_date(week):
continue
l.append(w)
for w in workers:
if w.non_working_day_list:
l = []
for (
nwd,
start_nwd_date,
end_nwd_date,
) in w.non_working_day_list:
if nwd in sb:
if start_nwd_date > self.get_week_start_date(week):
continue
if end_nwd_date < self.get_week_start_date(week):
continue
l.append(w)
if l:
nwd_workers.append(w)
if l:
nwd_workers.append(w)
else:
full_workers.append(w)
else:
full_workers.append(w)
if not nwd_workers: # Just do the usual
self.model.constraints.add(
sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name, sb_idx
]
for worker in self.workers
if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned
)
<= workers_required
)
else:
full_workers.append(w)
if not nwd_workers: # Just do the usual
self.model.constraints.add(
sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name
]
for worker in self.workers
self.model.constraints.add(
sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name, sb_idx
]
for worker in full_workers
if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned
)
<= workers_required + 1
)
<= workers_required
)
else:
self.model.constraints.add(
sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name
]
for worker in full_workers
elif shift.force_as_block:
if workers_required:
self.model.constraints.add(
sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name, sb_idx
]
for worker in self.workers
if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned
)
<= workers_required
)
<= workers_required + 1
)
elif shift.force_as_block:
if workers_required:
self.model.constraints.add(
sum(
self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name
]
for worker in self.workers
)
<= workers_required
)
# Most of our constraints apply per worker
# Worker constraint loop (worker loop)
@@ -4185,24 +4197,27 @@ class RotaBuilder(object):
if self.constraint_options_model.balance_blocks:
blocks_balancing = sum(
block_shift_balancing_constant
* self.model.blocks_assigned[week, shift]
* self.model.blocks_assigned[week, shift, sb_idx]
for week in self.weeks
for shift in self.shifts_to_assign_as_blocks()
for sb_idx in range(len(self.get_shift_sub_blocks(self.get_shift_by_name(shift))))
if (week, shift, sb_idx) in self.model.blocks_assigned
)
else:
blocks_balancing = 0
prefer_block_expr = sum(
worker.assign_as_block_preferences.get(shift.name, 0)
* self.model.blocks_worker_shift_assigned[worker.id, week, shift.name]
* self.model.blocks_worker_shift_assigned[worker.id, week, shift.name, sb_idx]
for worker in self.workers
for week in self.weeks
for shift in self.shifts
if hasattr(worker, "assign_as_block_preferences")
and shift.name in worker.assign_as_block_preferences
and (worker.id, week, shift.name)
in self.model.blocks_worker_shift_assigned
and self.is_worker_eligible_for_shift(worker, shift)
for sb_idx in range(len(self.get_shift_sub_blocks(shift)))
if (worker.id, week, shift.name, sb_idx)
in self.model.blocks_worker_shift_assigned
)
# Quadratic :(
@@ -5046,6 +5061,27 @@ class RotaBuilder(object):
)
])
def get_shift_sub_blocks(self, shift: SingleShift) -> List[List[str]]:
if getattr(shift, "force_as_block_unless_nwd", None):
val = shift.force_as_block_unless_nwd
if isinstance(val, (list, tuple)) and all(isinstance(x, (list, tuple, set)) for x in val):
return [list(x) for x in val]
return [list(shift.days)]
# Check if it is a block shift globally or per-worker
is_block = (
shift.assign_as_block
or shift.force_as_block
or any(
hasattr(worker, "assign_as_block_preferences")
and getattr(worker, "assign_as_block_preferences", {}).get(shift.name, 0) != 0
for worker in self.workers
)
)
if is_block:
return [list(shift.days)]
return []
def get_all_locum_availability(self):
return self.locum_availability_map