Compare commits

..
6 Commits
9 changed files with 631 additions and 111 deletions
+11 -1
View File
@@ -16,10 +16,13 @@ try:
WorkRequests, WorkRequests,
HardDayDependency, HardDayDependency,
HardDayExclusion, HardDayExclusion,
ShiftStartDate,
ShiftEndDate,
) )
except Exception: except Exception:
NonWorkingDays = OutOfProgramme = NotAvailableToWork = PreferenceNotToWork = None NonWorkingDays = OutOfProgramme = NotAvailableToWork = PreferenceNotToWork = None
WorkRequests = HardDayDependency = HardDayExclusion = None WorkRequests = HardDayDependency = HardDayExclusion = None
ShiftStartDate = ShiftEndDate = None
class WorkerForm(forms.ModelForm): class WorkerForm(forms.ModelForm):
@@ -46,6 +49,8 @@ class WorkerForm(forms.ModelForm):
assign_as_block_preferences = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->weight", label="Block preferences (JSON)") assign_as_block_preferences = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->weight", label="Block preferences (JSON)")
shift_fte_overrides = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->fte", label="Shift FTE overrides (JSON)") shift_fte_overrides = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->fte", label="Shift FTE overrides (JSON)")
exact_shifts = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->exact_count", label="Exact shifts (JSON)") exact_shifts = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->exact_count", label="Exact shifts (JSON)")
shift_start_dates = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->start_date (YYYY-MM-DD)", label="Shift start dates (JSON)")
shift_end_dates = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON mapping shift_name->end_date (YYYY-MM-DD)", label="Shift end dates (JSON)")
groups = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON list of group names (e.g. [\"group1\", \"group2\"])", label="Groups (JSON)") groups = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON list of group names (e.g. [\"group1\", \"group2\"])", label="Groups (JSON)")
previous_shifts = forms.CharField(required=False, widget=forms.Textarea, help_text="Free JSON structure for previous shifts", label="Previous shifts (JSON)") previous_shifts = forms.CharField(required=False, widget=forms.Textarea, help_text="Free JSON structure for previous shifts", label="Previous shifts (JSON)")
shift_balance_extra = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON structure for extra shift-balance penalties", label="Shift balance extra (JSON)") shift_balance_extra = forms.CharField(required=False, widget=forms.Textarea, help_text="JSON structure for extra shift-balance penalties", label="Shift balance extra (JSON)")
@@ -97,6 +102,8 @@ class WorkerForm(forms.ModelForm):
"assign_as_block_preferences", "assign_as_block_preferences",
"shift_fte_overrides", "shift_fte_overrides",
"exact_shifts", "exact_shifts",
"shift_start_dates",
"shift_end_dates",
"groups", "groups",
"previous_shifts", "previous_shifts",
"shift_balance_extra", "shift_balance_extra",
@@ -151,6 +158,8 @@ class WorkerForm(forms.ModelForm):
"assign_as_block_preferences", "assign_as_block_preferences",
"shift_fte_overrides", "shift_fte_overrides",
"exact_shifts", "exact_shifts",
"shift_start_dates",
"shift_end_dates",
"groups", "groups",
"previous_shifts", "previous_shifts",
"shift_balance_extra", "shift_balance_extra",
@@ -272,9 +281,10 @@ class WorkerForm(forms.ModelForm):
cleaned["work_requests"] = _parse_and_validate("work_requests", WorkRequests) cleaned["work_requests"] = _parse_and_validate("work_requests", WorkRequests)
cleaned["locum_availability"] = _parse_and_validate("locum_availability", WorkRequests) cleaned["locum_availability"] = _parse_and_validate("locum_availability", WorkRequests)
# Hard day dependency/exclusion structures
cleaned["hard_day_dependencies"] = _parse_and_validate("hard_day_dependencies", HardDayDependency) cleaned["hard_day_dependencies"] = _parse_and_validate("hard_day_dependencies", HardDayDependency)
cleaned["hard_day_exclusions"] = _parse_and_validate("hard_day_exclusions", HardDayExclusion) cleaned["hard_day_exclusions"] = _parse_and_validate("hard_day_exclusions", HardDayExclusion)
cleaned["shift_start_dates"] = _parse_and_validate("shift_start_dates", ShiftStartDate)
cleaned["shift_end_dates"] = _parse_and_validate("shift_end_dates", ShiftEndDate)
# forced assignments: expect list of tuples [week, day, shift] # forced assignments: expect list of tuples [week, day, shift]
fa = cleaned.get("forced_assignments") fa = cleaned.get("forced_assignments")
+2
View File
@@ -242,6 +242,8 @@ class Worker(models.Model):
"hard_day_exclusions", "hard_day_exclusions",
"forced_assignments", "forced_assignments",
"forced_assignments_by_date", "forced_assignments_by_date",
"shift_start_dates",
"shift_end_dates",
] ]
for k in int_keys: for k in int_keys:
+26 -3
View File
@@ -507,11 +507,34 @@ def load_workers():
def parse_time_limit(t) -> int:
"""Parses a time limit string (e.g. '10h', '30m', '45s') or an integer to seconds."""
if isinstance(t, int):
return t
if isinstance(t, float):
return int(t)
t = str(t).strip().lower()
if not t:
return 0
if t.isdigit():
return int(t)
if t.endswith("h"):
return int(float(t[:-1]) * 3600)
if t.endswith("m"):
return int(float(t[:-1]) * 60)
if t.endswith("s"):
return int(float(t[:-1]))
try:
return int(float(t))
except ValueError:
raise ValueError(f"Invalid time format: {t}. Expected format like '10h', '30m', '45s' or a number of seconds.")
@app.command() @app.command()
def main( def main(
suspend: bool = False, suspend: bool = False,
solve: bool = True, solve: bool = True,
time_to_run: int = 60 * 60, time_to_run: str = "1h",
ratio: float = 0.1, ratio: float = 0.1,
start_date: datetime.datetime = ROTA_START_DATE, start_date: datetime.datetime = ROTA_START_DATE,
weeks: int = 22, weeks: int = 22,
@@ -698,8 +721,8 @@ def main(
# Rota.build_workers() # Rota.build_workers()
# Rota.build_model() # Rota.build_model()
solver_options = {"ratio": ratio, "seconds": time_to_run, "threads": 10} seconds_to_run = parse_time_limit(time_to_run)
# solver_options = {"seconds": time_to_run, "threads": 10} solver_options = {"ratio": ratio, "seconds": seconds_to_run, "threads": 10}
# start_time = time.time() # start_time = time.time()
Rota.build_and_solve(solver_options, export=True, export_with_timestamp=False, solve=solve, solver="appsi_highs") Rota.build_and_solve(solver_options, export=True, export_with_timestamp=False, solve=solve, solver="appsi_highs")
+49 -6
View File
@@ -34,6 +34,7 @@ sites = (
"truro ir", "truro ir",
"exeter ir", "exeter ir",
"torbay ir", "torbay ir",
"ir nights",
) )
from rota_generator.workers import ( from rota_generator.workers import (
@@ -45,12 +46,34 @@ from rota_generator.workers import (
OutOfProgramme, OutOfProgramme,
) )
def parse_time_limit(t) -> int:
"""Parses a time limit string (e.g. '10h', '30m', '45s') or an integer to seconds."""
if isinstance(t, int):
return t
if isinstance(t, float):
return int(t)
t = str(t).strip().lower()
if not t:
return 0
if t.isdigit():
return int(t)
if t.endswith("h"):
return int(float(t[:-1]) * 3600)
if t.endswith("m"):
return int(float(t[:-1]) * 60)
if t.endswith("s"):
return int(float(t[:-1]))
try:
return int(float(t))
except ValueError:
raise ValueError(f"Invalid time format: {t}. Expected format like '10h', '30m', '45s' or a number of seconds.")
@app.command() @app.command()
def main( def main(
suspend: bool = False, suspend: bool = False,
solve: bool = True, solve: bool = True,
time_to_run: int = 60 * 60 * 10, time_to_run: str = "10h",
ratio: float = 0.1, ratio: float = 0.1,
start_date: datetime.datetime = "2026-09-07", start_date: datetime.datetime = "2026-09-07",
weeks: int = 26, weeks: int = 26,
@@ -74,6 +97,7 @@ def main(
Rota.constraint_options["max_weekend_frequency"] = 3 Rota.constraint_options["max_weekend_frequency"] = 3
Rota.constraint_options["max_days_per_week_block"] = [(8,3)] Rota.constraint_options["max_days_per_week_block"] = [(8,3)]
Rota.constraint_options["maximum_allowed_shift_diff"] = 1
# wr = [ # wr = [
# WorkerRequirement( # WorkerRequirement(
@@ -174,11 +198,11 @@ def main(
length=12.5, length=12.5,
days=days[4:], days=days[4:],
# balance_offset=3, # balance_offset=3,
rota_on_nwds=True, #rota_on_nwds=True,
force_as_block=True, #force_as_block=True,
# assign_as_block=True, # assign_as_block=True,
constraints=[PreShiftConstraint(days=2), PostShiftConstraint(days=2)], constraints=[PreShiftConstraint(days=2), PostShiftConstraint(days=2)],
# force_as_block_unless_nwd=True force_as_block_unless_nwd=[["Fri"], ["Sat", "Sun"]],
), ),
SingleShift( SingleShift(
sites=("torbay", "torbay twilights and weekends"), sites=("torbay", "torbay twilights and weekends"),
@@ -503,6 +527,24 @@ def main(
"night_weekday": 0, "night_weekday": 0,
} }
override_shift_start_dates = []
if worker in ["Anushka Kulkarni"]:
override_shift_start_dates = [
{
"shift": "night_weekday",
"start_date": datetime.datetime.strptime(
"2026-11-01", "%Y-%m-%d"
).date(),
},
{
"shift": "night_weekend",
"start_date": datetime.datetime.strptime(
"2026-11-01", "%Y-%m-%d"
).date(),
},
]
# if worker_name == "Hadi Mohamed": # if worker_name == "Hadi Mohamed":
# shift_fte_overrides = { # shift_fte_overrides = {
# "weekend_exeter": 50, # "weekend_exeter": 50,
@@ -536,6 +578,7 @@ def main(
bank_holiday_extra=w["bank_holiday_extra"], bank_holiday_extra=w["bank_holiday_extra"],
shift_fte_overrides=shift_fte_overrides, shift_fte_overrides=shift_fte_overrides,
exact_shifts=exact_shifts, exact_shifts=exact_shifts,
shift_start_dates=override_shift_start_dates,
) )
# print(w) # print(w)
@@ -546,8 +589,8 @@ def main(
# Rota.build_workers() # Rota.build_workers()
# Rota.build_model() # Rota.build_model()
solver_options = {"ratio": ratio, "seconds": time_to_run, "threads": 10} seconds_to_run = parse_time_limit(time_to_run)
# solver_options = {"seconds": time_to_run, "threads": 10} solver_options = {"ratio": ratio, "seconds": seconds_to_run, "threads": 10}
# start_time = time.time() # start_time = time.time()
Rota.build_and_solve(solver_options, export=True, solve=solve, solver="appsi_highs", export_with_timestamp=True) Rota.build_and_solve(solver_options, export=True, solve=solve, solver="appsi_highs", export_with_timestamp=True)
+2 -2
View File
@@ -33,7 +33,7 @@ def load_leave(Rota):
if live_rota: if live_rota:
download = s.get( download = s.get(
"https://docs.google.com/spreadsheets/d/e/2PACX-1vSh64r9uUbPsCdhcV7zmZz10gVqIy6rhXFUvDgJxRD5W45IYEaZpKMxAqDsD1ph6dob4AiRzJXVVtOu/pub?gid=580824586&single=true&output=csv", "https://docs.google.com/spreadsheets/d/e/2PACX-1vSh64r9uUbPsCdhcV7zmZz10gVqIy6rhXFUvDgJxRD5W45IYEaZpKMxAqDsD1ph6dob4AiRzJXVVtOu/pub?gid=396859373&single=true&output=csv",
headers=headers, headers=headers,
) )
# download = s.get("https://docs.google.com/spreadsheets/d/e/2PACX-1vSRx9VWXSlRubPyA0RhiI-Oqf5eHNYYEc6rFzlraDbR5_8qqr5g13-4uV-gn4u-TjZxiSMv1fBUaESq/pub?gid=814517272&single=true&output=csv", headers=headers) # download = s.get("https://docs.google.com/spreadsheets/d/e/2PACX-1vSRx9VWXSlRubPyA0RhiI-Oqf5eHNYYEc6rFzlraDbR5_8qqr5g13-4uV-gn4u-TjZxiSMv1fBUaESq/pub?gid=814517272&single=true&output=csv", headers=headers)
@@ -295,7 +295,7 @@ def load_academy(Rota):
""" """
# Use the same published CSV as load_leave (main sheet) # Use the same published CSV as load_leave (main sheet)
url = ( url = (
"https://docs.google.com/spreadsheets/d/e/2PACX-1vSh64r9uUbPsCdhcV7zmZz10gVqIy6rhXFUvDgJxRD5W45IYEaZpKMxAqDsD1ph6dob4AiRzJXVVtOu/pub?gid=580824586&single=true&output=csv" "https://docs.google.com/spreadsheets/d/e/2PACX-1vSh64r9uUbPsCdhcV7zmZz10gVqIy6rhXFUvDgJxRD5W45IYEaZpKMxAqDsD1ph6dob4AiRzJXVVtOu/pub?gid=396859373&single=true&output=csv"
) )
with Session() as s: with Session() as s:
+110 -26
View File
@@ -417,7 +417,7 @@ class SingleShift(BaseModel):
rota_on_nwds: bool = False rota_on_nwds: bool = False
assign_as_block: bool = False assign_as_block: bool = False
force_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 hard_constrain_shift: bool = True
bank_holidays_only: bool = False bank_holidays_only: bool = False
#constraint: list[ShiftConstraint] = [] #constraint: list[ShiftConstraint] = []
@@ -1050,20 +1050,28 @@ class RotaBuilder(object):
initialize=0, 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( self.model.blocks_worker_shift_assigned = Var(
( blocks_worker_keys,
(worker.id, week, shift)
for worker, week, shift in self.worker_week_shifts_to_assign_as_blocks()
),
within=Binary, within=Binary,
initialize=0, 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( self.model.blocks_assigned = Var(
( blocks_assigned_keys,
(week, shift)
for week, shift in self.week_shifts_to_assign_as_blocks()
),
within=NonNegativeIntegers, within=NonNegativeIntegers,
initialize=0, initialize=0,
) )
@@ -1900,30 +1908,33 @@ class RotaBuilder(object):
for week, shift_name in self.week_shifts_to_assign_as_blocks(): for week, shift_name in self.week_shifts_to_assign_as_blocks():
shift = self.get_shift_by_name(shift_name) shift = self.get_shift_by_name(shift_name)
sub_blocks = self.get_shift_sub_blocks(shift)
for sb_idx, sb in enumerate(sub_blocks):
for worker in self.get_workers_for_shift(shift): for worker in self.get_workers_for_shift(shift):
if (worker.id, week, shift.name) in self.model.blocks_worker_shift_assigned: if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned:
try: try:
self.model.constraints.add( self.model.constraints.add(
8 8
* self.model.blocks_worker_shift_assigned[ * self.model.blocks_worker_shift_assigned[
worker.id, week, shift_name worker.id, week, shift_name, sb_idx
] ]
>= sum( >= sum(
self.model.works[worker.id, week, day, shift_name] self.model.works[worker.id, week, day, shift_name]
for day in shift.days for day in sb
) )
) )
except KeyError: except KeyError:
pass pass
self.model.constraints.add( self.model.constraints.add(
self.model.blocks_assigned[week, shift_name] self.model.blocks_assigned[week, shift_name, sb_idx]
== sum( == sum(
self.model.blocks_worker_shift_assigned[ self.model.blocks_worker_shift_assigned[
worker.id, week, shift_name worker.id, week, shift_name, sb_idx
] ]
for worker in self.workers for worker in self.workers
if (worker.id, week, shift.name) in self.model.blocks_worker_shift_assigned if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned
) )
) )
@@ -1933,7 +1944,7 @@ class RotaBuilder(object):
if shift.force_as_block_unless_nwd: if shift.force_as_block_unless_nwd:
workers = self.get_workers_for_shift(shift) workers = self.get_workers_for_shift(shift)
# Get workers who have a nwd on the shift # Get workers who have a nwd on the sub-block
nwd_workers = [] nwd_workers = []
full_workers = [] full_workers = []
@@ -1945,7 +1956,7 @@ class RotaBuilder(object):
start_nwd_date, start_nwd_date,
end_nwd_date, end_nwd_date,
) in w.non_working_day_list: ) in w.non_working_day_list:
if nwd in shift.days: if nwd in sb:
if start_nwd_date > self.get_week_start_date(week): if start_nwd_date > self.get_week_start_date(week):
continue continue
if end_nwd_date < self.get_week_start_date(week): if end_nwd_date < self.get_week_start_date(week):
@@ -1956,7 +1967,6 @@ class RotaBuilder(object):
nwd_workers.append(w) nwd_workers.append(w)
else: else:
full_workers.append(w) full_workers.append(w)
else: else:
full_workers.append(w) full_workers.append(w)
@@ -1964,20 +1974,21 @@ class RotaBuilder(object):
self.model.constraints.add( self.model.constraints.add(
sum( sum(
self.model.blocks_worker_shift_assigned[ self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name worker.id, week, shift.name, sb_idx
] ]
for worker in self.workers for worker in self.workers
if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned
) )
<= workers_required <= workers_required
) )
else: else:
self.model.constraints.add( self.model.constraints.add(
sum( sum(
self.model.blocks_worker_shift_assigned[ self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name worker.id, week, shift.name, sb_idx
] ]
for worker in full_workers 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 + 1
) )
@@ -1987,9 +1998,10 @@ class RotaBuilder(object):
self.model.constraints.add( self.model.constraints.add(
sum( sum(
self.model.blocks_worker_shift_assigned[ self.model.blocks_worker_shift_assigned[
worker.id, week, shift.name worker.id, week, shift.name, sb_idx
] ]
for worker in self.workers for worker in self.workers
if (worker.id, week, shift.name, sb_idx) in self.model.blocks_worker_shift_assigned
) )
<= workers_required <= workers_required
) )
@@ -3726,6 +3738,22 @@ class RotaBuilder(object):
) )
) )
# Shift-specific active dates constraint
for shift in self.get_shifts():
if (worker.id, week, day, shift.name) in self.model.works:
calc_start = getattr(worker, "calculated_shift_start_dates", {}).get(shift.name, worker.calculated_start_date)
calc_end = getattr(worker, "calculated_shift_end_dates", {}).get(shift.name, worker.calculated_end_date)
if calc_start is not None and calc_end is not None:
date = self.week_day_date_map[(week, day)]
if date < calc_start or date >= calc_end:
self.model.constraints.add(
self.model.works[worker.id, week, day, shift.name] == 0
)
if self.get_locum_workers() and (worker.id, week, day, shift.name) in self.model.locum_works:
self.model.constraints.add(
self.model.locum_works[worker.id, week, day, shift.name] == 0
)
# single shift per day (unless multi-shift allowed) # single shift per day (unless multi-shift allowed)
# This is signifantly slower so only enable if required # This is signifantly slower so only enable if required
shifts_today = self.get_shift_names_by_week_day(week, day) shifts_today = self.get_shift_names_by_week_day(week, day)
@@ -4185,24 +4213,27 @@ class RotaBuilder(object):
if self.constraint_options_model.balance_blocks: if self.constraint_options_model.balance_blocks:
blocks_balancing = sum( blocks_balancing = sum(
block_shift_balancing_constant block_shift_balancing_constant
* self.model.blocks_assigned[week, shift] * self.model.blocks_assigned[week, shift, sb_idx]
for week in self.weeks for week in self.weeks
for shift in self.shifts_to_assign_as_blocks() 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: else:
blocks_balancing = 0 blocks_balancing = 0
prefer_block_expr = sum( prefer_block_expr = sum(
worker.assign_as_block_preferences.get(shift.name, 0) 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 worker in self.workers
for week in self.weeks for week in self.weeks
for shift in self.shifts for shift in self.shifts
if hasattr(worker, "assign_as_block_preferences") if hasattr(worker, "assign_as_block_preferences")
and shift.name in 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) 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 :( # Quadratic :(
@@ -4308,6 +4339,38 @@ class RotaBuilder(object):
"Invalid exact shift", "Invalid exact shift",
f"Worker {worker.name} requested exact shifts for non-existent shift {s_name}", f"Worker {worker.name} requested exact shifts for non-existent shift {s_name}",
) )
# Validate shift-dependent start/end dates
start_dates = getattr(worker, "shift_start_dates", [])
end_dates = getattr(worker, "shift_end_dates", [])
start_dict = {}
if isinstance(start_dates, dict):
start_dict = start_dates
elif isinstance(start_dates, list):
for item in start_dates:
if hasattr(item, "shift"):
start_dict[item.shift] = item.start_date
end_dict = {}
if isinstance(end_dates, dict):
end_dict = end_dates
elif isinstance(end_dates, list):
for item in end_dates:
if hasattr(item, "shift"):
end_dict[item.shift] = item.end_date
for s_name in set(list(start_dict.keys()) + list(end_dict.keys())):
if s_name not in self.shifts_by_name:
raise InvalidShift(
f"Worker {worker.name} specified shift-dependent dates for non-existent shift {s_name}"
)
s = self.get_shift_by_name(s_name)
if not self.is_worker_eligible_for_shift(worker, s):
self.add_warning(
"Worker/ineligible shift date constraint",
f"Worker {worker.name} specified shift-dependent dates for shift {s_name} but is not eligible to work it"
)
wid = worker.id wid = worker.id
if wid in self.workers_id_map: if wid in self.workers_id_map:
message = f"Worker with id '{wid}' has been added twice" message = f"Worker with id '{wid}' has been added twice"
@@ -5046,6 +5109,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): def get_all_locum_availability(self):
return self.locum_availability_map return self.locum_availability_map
+158 -3
View File
@@ -169,6 +169,42 @@ class MaxUniqueShiftsPerWeekBlockConstraint(BaseModel):
) )
class ShiftStartDate(BaseModel):
shift: str
start_date: datetime.date
@field_validator("start_date", mode="before")
@classmethod
def coerce_date(cls, v):
if isinstance(v, datetime.date):
return v
if isinstance(v, str):
for fmt in ("%d/%m/%Y", "%d/%m/%y", "%Y-%m-%d"):
try:
return datetime.datetime.strptime(v, fmt).date()
except Exception:
continue
raise ValueError(f"Cannot parse date: {v}")
class ShiftEndDate(BaseModel):
shift: str
end_date: datetime.date
@field_validator("end_date", mode="before")
@classmethod
def coerce_date(cls, v):
if isinstance(v, datetime.date):
return v
if isinstance(v, str):
for fmt in ("%d/%m/%Y", "%d/%m/%y", "%Y-%m-%d"):
try:
return datetime.datetime.strptime(v, fmt).date()
except Exception:
continue
raise ValueError(f"Cannot parse date: {v}")
class Worker(BaseModel): class Worker(BaseModel):
name: str name: str
site: str site: str
@@ -215,11 +251,55 @@ class Worker(BaseModel):
exact_shifts: dict[str, int] = {} # Map shift_name to exact number of shifts exact_shifts: dict[str, int] = {} # Map shift_name to exact number of shifts
groups: set[str] = set() # Groups that the worker belongs to groups: set[str] = set() # Groups that the worker belongs to
shift_start_dates: list[ShiftStartDate] = [] # List of ShiftStartDate models
shift_end_dates: list[ShiftEndDate] = [] # List of ShiftEndDate models
weekend_shift_target_number: float = 0.0 weekend_shift_target_number: float = 0.0
avoid_shifts_on_dates: list[AvoidShiftOnDates] = [] avoid_shifts_on_dates: list[AvoidShiftOnDates] = []
@field_validator("shift_start_dates", mode="before")
@classmethod
def coerce_shift_start_dates(cls, v):
if v is None:
return []
if isinstance(v, dict):
res = []
for shift, d in v.items():
res.append(ShiftStartDate(shift=shift, start_date=d))
return res
if isinstance(v, list):
res = []
for item in v:
if isinstance(item, ShiftStartDate):
res.append(item)
elif isinstance(item, dict):
res.append(ShiftStartDate(**item))
else:
raise ValueError(f"Invalid ShiftStartDate item: {item}")
return res
raise ValueError("Must be a list or dictionary")
@field_validator("shift_end_dates", mode="before")
@classmethod
def coerce_shift_end_dates(cls, v):
if v is None:
return []
if isinstance(v, dict):
res = []
for shift, d in v.items():
res.append(ShiftEndDate(shift=shift, end_date=d))
return res
if isinstance(v, list):
res = []
for item in v:
if isinstance(item, ShiftEndDate):
res.append(item)
elif isinstance(item, dict):
res.append(ShiftEndDate(**item))
else:
raise ValueError(f"Invalid ShiftEndDate item: {item}")
return res
raise ValueError("Must be a list or dictionary")
model_config = ConfigDict( model_config = ConfigDict(
extra="allow", extra="allow",
@@ -416,13 +496,88 @@ class Worker(BaseModel):
self.proportion_rota_to_work = days_to_work / Rota.rota_days_length self.proportion_rota_to_work = days_to_work / Rota.rota_days_length
self.days_to_work = days_to_work self.days_to_work = days_to_work
# Shift-specific start/end dates and adjusted FTEs
self.calculated_shift_start_dates = {}
self.calculated_shift_end_dates = {}
self.proportion_rota_to_work_shifts = {}
# Convert list of ShiftStartDate/ShiftEndDate models/dicts to dictionary mapping for quick lookup
start_dates_dict = {item.shift: item.start_date for item in getattr(self, "shift_start_dates", []) if hasattr(item, "shift")}
end_dates_dict = {item.shift: item.end_date for item in getattr(self, "shift_end_dates", []) if hasattr(item, "shift")}
for shift in Rota.get_shifts():
s_date = start_dates_dict.get(shift.name, self.start_date)
if s_date is None:
calc_s = self.calculated_start_date
else:
if isinstance(s_date, str):
for fmt in ("%d/%m/%Y", "%d/%m/%y", "%Y-%m-%d"):
try:
s_date = datetime.datetime.strptime(s_date, fmt).date()
break
except Exception:
continue
calc_s = s_date
if calc_s < Rota.start_date:
calc_s = Rota.start_date
elif calc_s > Rota.rota_end_date:
calc_s = Rota.rota_end_date
e_date = end_dates_dict.get(shift.name, self.end_date)
if e_date is None:
calc_e = self.calculated_end_date
else:
if isinstance(e_date, str):
for fmt in ("%d/%m/%Y", "%d/%m/%y", "%Y-%m-%d"):
try:
e_date = datetime.datetime.strptime(e_date, fmt).date()
break
except Exception:
continue
calc_e = e_date
if calc_e > Rota.rota_end_date:
calc_e = Rota.rota_end_date
elif calc_e < Rota.start_date:
calc_e = Rota.start_date
if calc_s >= calc_e:
calc_s = Rota.rota_end_date
calc_e = Rota.rota_end_date
self.calculated_shift_start_dates[shift.name] = calc_s
self.calculated_shift_end_dates[shift.name] = calc_e
days_active = (calc_e - calc_s).days
# Subtract overlapping OOP days from the active period of this shift
for item in self.oop:
start_oop = item.start_date
end_oop = item.end_date
if isinstance(start_oop, datetime.date):
start_oop_date = start_oop
else:
start_oop_date = datetime.datetime.strptime(start_oop, "%d/%m/%y").date()
if isinstance(end_oop, datetime.date):
end_oop_date = end_oop
else:
end_oop_date = datetime.datetime.strptime(end_oop, "%d/%m/%y").date()
overlap_start = max(calc_s, start_oop_date)
overlap_end = min(calc_e, end_oop_date)
if overlap_start < overlap_end:
days_active -= (overlap_end - overlap_start).days
self.proportion_rota_to_work_shifts[shift.name] = max(0.0, days_active / Rota.rota_days_length)
# We have to adjust the full time equivalent for people who CCT / leave the rota early # We have to adjust the full time equivalent for people who CCT / leave the rota early
self.fte_adj = self.fte * self.proportion_rota_to_work self.fte_adj = self.fte * self.proportion_rota_to_work
self.fte_adj_shifts = {} self.fte_adj_shifts = {}
if self.shift_fte_overrides: for shift in Rota.get_shifts():
for shift, fte in self.shift_fte_overrides.items(): prop = self.proportion_rota_to_work_shifts.get(shift.name, self.proportion_rota_to_work)
self.fte_adj_shifts[shift] = fte * self.proportion_rota_to_work base_fte = self.shift_fte_overrides.get(shift.name, self.fte)
self.fte_adj_shifts[shift.name] = base_fte * prop
if self.fte_adj > 100: if self.fte_adj > 100:
+58 -1
View File
@@ -163,7 +163,7 @@ def test_nwd_force_as_block_force_split():
for worker in Rota.workers: for worker in Rota.workers:
shifts = Rota.get_worker_shift_list(worker) shifts = Rota.get_worker_shift_list(worker)
shifts_string = "".join([i if i != "" else "-" for i in shifts]) shifts_string = "".join([i if i != "" else "-" for i in shifts])
for week in weeks_from_list(shifts_string[7 * 5 :]): for week in weeks_from_list(shifts_string[7 * 6 :]):
assert week in ("-----ww", "dddddww") assert week in ("-----ww", "dddddww")
if worker.name == "worker1": if worker.name == "worker1":
assert shifts_string[: 7 * 5].count("ww-") == 4 assert shifts_string[: 7 * 5].count("ww-") == 4
@@ -214,3 +214,60 @@ def test_nwd_testing():
assert week == "dddddww" assert week == "dddddww"
else: else:
assert week == "-----ww" assert week == "-----ww"
def test_force_as_block_unless_nwd_with_subblocks():
# Test that a Fri/Sat/Sun shift can be split into Fri and Sat/Sun blocks
# if a worker does not work on Fridays.
Rota = setup_rota(weeks_to_rota=2)
start_date = Rota.start_date
# worker1 does not work on Fridays (Fri is NWD)
worker1 = Worker(
name="worker1", site="group1", grade=1,
nwds=[{"day": "Fri", "start_date": start_date, "end_date": start_date + datetime.timedelta(weeks=2)}],
)
# worker2 works normally
worker2 = Worker(name="worker2", site="group1", grade=1)
Rota.add_workers((worker1, worker2))
# Fri/Sat/Sun shift, requiring 1 worker, forced as block unless NWD
# We specify sub-blocks: [["Fri"], ["Sat", "Sun"]]
Rota.add_shifts(
SingleShift(
sites=("group1",), name="weekend", length=12.5,
days=["Fri", "Sat", "Sun"],
workers_required=1,
force_as_block_unless_nwd=[["Fri"], ["Sat", "Sun"]],
),
)
# We remove no valid shifts warning
Rota.terminate_on_warning.remove("Worker/no valid shifts")
Rota.build_and_solve()
assert Rota.results.solver.status == "ok"
# Check assignments:
# Since worker1 has NWD on Friday:
# - worker1 cannot work Friday.
# - But worker1 can work Sat and Sun as a block!
# So on Saturday and Sunday, worker1 should be assigned.
# On Friday, worker2 should be assigned.
for week in Rota.weeks:
w1_fri = Rota.model.works[worker1.id, week, "Fri", "weekend"].value
w1_sat = Rota.model.works[worker1.id, week, "Sat", "weekend"].value
w1_sun = Rota.model.works[worker1.id, week, "Sun", "weekend"].value
w2_fri = Rota.model.works[worker2.id, week, "Fri", "weekend"].value
w2_sat = Rota.model.works[worker2.id, week, "Sat", "weekend"].value
w2_sun = Rota.model.works[worker2.id, week, "Sun", "weekend"].value
assert w1_fri == 0
# Either worker1 works Sat/Sun and worker2 works Fri (split block)
# OR worker2 works Fri/Sat/Sun (full block) and worker1 works nothing
is_split = (w1_sat == 1 and w1_sun == 1 and w2_fri == 1 and w2_sat == 0 and w2_sun == 0)
is_full_w2 = (w1_sat == 0 and w1_sun == 0 and w2_fri == 1 and w2_sat == 1 and w2_sun == 1)
assert is_split or is_full_w2
+147 -1
View File
@@ -3,7 +3,7 @@ from rota_generator.shifts import InvalidShift, MaxShiftsPerWeekConstraint, NoWo
from pydantic import ValidationError from pydantic import ValidationError
import datetime import datetime
from rota_generator.workers import NonWorkingDays, NotAvailableToWork, Worker, generate_not_available_to_works from rota_generator.workers import NonWorkingDays, NotAvailableToWork, Worker, generate_not_available_to_works, ShiftStartDate, ShiftEndDate
from rota_generator.workers import WorkRequests from rota_generator.workers import WorkRequests
from rota_generator.shifts import PreShiftConstraint from rota_generator.shifts import PreShiftConstraint
@@ -1985,3 +1985,149 @@ def test_shift_group_availability_exclusion():
assert worker_shift_counts["worker1"] == 2 assert worker_shift_counts["worker1"] == 2
assert worker_shift_counts["worker2"] == 0 assert worker_shift_counts["worker2"] == 0
def test_shift_dependent_active_dates_constraints():
weeks_to_rota = 10
start_date = datetime.date(2022, 3, 7) # Mon
Rota = RotaBuilder(
start_date,
weeks_to_rota=weeks_to_rota,
)
worker1 = Worker(name="worker1", site="group1", grade=1, fte=100)
worker2 = Worker(
name="worker2", site="group1", grade=1, fte=100,
shift_start_dates={"a": datetime.date(2022, 3, 28)},
shift_end_dates={"a": datetime.date(2022, 4, 25)},
)
Rota.add_workers([worker1, worker2])
Rota.add_shifts(
SingleShift(
sites=("group1",),
name="a",
length=12.5,
days=("Mon",),
workers_required=1,
),
)
Rota.build_and_solve(options={"ratio": 0.0})
assert Rota.results.solver.status == "ok"
# Verify assignments for worker2
for week in range(1, weeks_to_rota + 1):
val = Rota.model.works[worker2.id, week, "Mon", "a"].value
assigned = val is not None and val > 0.5
if week < 4 or week >= 8:
assert not assigned, f"Worker 2 should not be assigned 'a' in week {week}"
def test_shift_dependent_fte_target_scaling():
weeks_to_rota = 10
start_date = datetime.date(2022, 3, 7) # Mon
Rota = RotaBuilder(
start_date,
weeks_to_rota=weeks_to_rota,
)
worker1 = Worker(name="worker1", site="group1", grade=1, fte=100)
worker2 = Worker(
name="worker2", site="group1", grade=1, fte=100,
shift_start_dates={"a": datetime.date(2022, 3, 28)},
shift_end_dates={"a": datetime.date(2022, 4, 25)},
)
Rota.add_workers([worker1, worker2])
Rota.add_shifts(
SingleShift(
sites=("group1",),
name="a",
length=12.5,
days=("Mon",),
workers_required=1,
),
)
Rota.build_and_solve(options={"ratio": 0.0})
assert Rota.results.solver.status == "ok"
# Targets check
assert abs(worker1.shift_target_number["a"] - 7.14) < 0.1
assert abs(worker2.shift_target_number["a"] - 2.86) < 0.1
def test_shift_dependent_dates_validation_invalid_shift():
weeks_to_rota = 2
start_date = datetime.date(2022, 3, 7)
Rota = RotaBuilder(start_date, weeks_to_rota=weeks_to_rota)
worker = Worker(
name="worker1", site="group1", grade=1, fte=100,
shift_start_dates=[ShiftStartDate(shift="non_existent", start_date=datetime.date(2022, 3, 7))]
)
Rota.add_worker(worker)
Rota.add_shifts(
SingleShift(
sites=("group1",),
name="a",
length=12.5,
days=("Mon",),
workers_required=1,
)
)
with pytest.raises(InvalidShift):
Rota.build_workers()
def test_shift_dependent_dates_validation_ineligible_warning():
weeks_to_rota = 2
start_date = datetime.date(2022, 3, 7)
Rota = RotaBuilder(start_date, weeks_to_rota=weeks_to_rota)
worker = Worker(
name="worker1", site="group1", grade=1, fte=100,
shift_start_dates=[ShiftStartDate(shift="a", start_date=datetime.date(2022, 3, 7))]
)
Rota.add_worker(worker)
Rota.add_shifts(
SingleShift(
sites=("group2",),
name="a",
length=12.5,
days=("Mon",),
workers_required=1,
)
)
Rota.terminate_on_warning.remove("Worker/no valid shifts")
Rota.build_workers()
warnings = Rota.get_warnings("Worker/ineligible shift date constraint")
assert len(warnings) == 1
assert "is not eligible to work it" in warnings[0][1]
def test_shift_dependent_dates_model_definition():
worker = Worker(
name="worker1", site="group1", grade=1, fte=100,
shift_start_dates=[{"shift": "a", "start_date": "2022-03-07"}],
shift_end_dates=[{"shift": "a", "end_date": "2022-04-07"}]
)
assert len(worker.shift_start_dates) == 1
assert isinstance(worker.shift_start_dates[0], ShiftStartDate)
assert worker.shift_start_dates[0].shift == "a"
assert worker.shift_start_dates[0].start_date == datetime.date(2022, 3, 7)
assert len(worker.shift_end_dates) == 1
assert isinstance(worker.shift_end_dates[0], ShiftEndDate)
assert worker.shift_end_dates[0].shift == "a"
assert worker.shift_end_dates[0].end_date == datetime.date(2022, 4, 7)
def test_parse_time_limit():
from gen_proc import parse_time_limit
assert parse_time_limit(3600) == 3600
assert parse_time_limit("3600") == 3600
assert parse_time_limit("10h") == 36000
assert parse_time_limit("30m") == 1800
assert parse_time_limit("45s") == 45
assert parse_time_limit(" 1.5h ") == 5400
with pytest.raises(ValueError):
parse_time_limit("invalid")