feat: add group management for workers and shifts, including validation for overlapping groups

This commit is contained in:
Ross
2026-07-09 22:14:30 +01:00
parent 33f50c7278
commit b5525bd3fb
6 changed files with 259 additions and 24 deletions
+91 -23
View File
@@ -346,7 +346,7 @@ def _pydantic_to_dict(obj):
# Recursively convert pydantic models, lists and dicts to JSON-serializable structures
if obj is None:
return None
if isinstance(obj, list):
if isinstance(obj, (list, set, tuple)):
return [_pydantic_to_dict(x) for x in obj]
if isinstance(obj, dict):
return {k: _pydantic_to_dict(v) for k, v in obj.items()}
@@ -433,6 +433,17 @@ class SingleShift(BaseModel):
pair_proxy: bool = False
display_char: str | None = None # Add this line
force_assign_with: list[str] = [] # List of shift names to force assign with this shift on the same day
include_groups: set[str] = set()
exclude_groups: set[str] = set()
@model_validator(mode="after")
def check_groups_non_overlapping(self) -> "SingleShift":
intersection = self.include_groups.intersection(self.exclude_groups)
if intersection:
raise ValueError(
f"include_groups and exclude_groups cannot overlap. Intersection: {intersection}"
)
return self
model_config = ConfigDict(
extra="allow",
@@ -855,6 +866,8 @@ class RotaBuilder(object):
options["time_limit"] = options.pop("seconds")
if "ratio" in options:
options["mip_rel_gap"] = options.pop("ratio")
if "threads" in options:
options.pop("threads")
else:
self.opt = SolverFactory(solver)
@@ -1610,13 +1623,16 @@ class RotaBuilder(object):
workers_required,
site_required,
) in self.get_required_workers_and_site_combinations():
# print(week, day, shift, workers_required, site_required)
shift_obj = self.get_shift_by_name(shift)
eligible_workers = [
worker for worker in self.workers
if self.is_worker_eligible_for_shift(worker, shift_obj)
]
self.model.constraints.add(
workers_required
== sum(
self.model.works[worker.id, week, day, shift]
for worker in self.workers
if worker.site in site_required
for worker in eligible_workers
)
)
@@ -1625,7 +1641,7 @@ class RotaBuilder(object):
self.model.locum_required[week, day, shift]
== sum(
self.model.works[worker.id, week, day, shift]
for worker in self.workers
for worker in eligible_workers
if worker.locum
)
)
@@ -1638,6 +1654,7 @@ class RotaBuilder(object):
# if not worker.locum
)
)
self.model.constraints.add(
0
== sum(
@@ -1647,14 +1664,17 @@ class RotaBuilder(object):
)
)
# And it is not assigned if the worker is from the wrong site
if [worker for worker in self.workers if worker.site not in site_required]:
# And it is not assigned if the worker is ineligible
ineligible_workers = [
worker for worker in self.workers
if not self.is_worker_eligible_for_shift(worker, shift_obj)
]
if ineligible_workers:
self.model.constraints.add(
0
== sum(
self.model.works[worker.id, week, day, shift]
for worker in self.workers
if worker.site not in site_required
for worker in ineligible_workers
)
)
# # Ensure shifts are only assigned by workers from the allowed sites
@@ -2647,14 +2667,14 @@ class RotaBuilder(object):
self.get_shifts(), description="Generate shift balance constraints"
):
if (
worker.site in shift.sites
self.is_worker_eligible_for_shift(worker, shift)
): # Each site specfies which sites self.workers can fullfill
total_shifts = self.shift_worker_counts[shift.name]
# Find workers who have exact shifts for this shift type
exact_workers = [
w for w in self.workers
if w.site in shift.sites and shift.name in getattr(w, "exact_shifts", {})
if self.is_worker_eligible_for_shift(w, shift) and shift.name in getattr(w, "exact_shifts", {})
]
exact_shifts_sum = sum(w.exact_shifts[shift.name] for w in exact_workers)
@@ -2671,7 +2691,7 @@ class RotaBuilder(object):
full_time_equivalent_joined = sum(
w.get_fte(shift=shift.name)
for w in self.workers
if w.site in shift.sites and w not in exact_workers
if self.is_worker_eligible_for_shift(w, shift) and w not in exact_workers
)
if not full_time_equivalent_joined:
@@ -3825,7 +3845,7 @@ class RotaBuilder(object):
continue
if day in constraint_shift.days:
# Only apply if worker can work this shift
if worker.site not in constraint_shift.sites:
if not self.is_worker_eligible_for_shift(worker, constraint_shift):
continue
try:
works = self.model.works[
@@ -3840,7 +3860,7 @@ class RotaBuilder(object):
pre_date = None
if pre_date is not None and pre_date in getattr(constraint, "exclude_dates", []):
continue
self.model.constraints.add(
1
>= works
@@ -3854,7 +3874,7 @@ class RotaBuilder(object):
)
if shiftname not in ignore_shifts
for w in workers
if w.site in constraint_shift.sites # Only workers who can work this shift
if self.is_worker_eligible_for_shift(w, constraint_shift) # Only workers who can work this shift
)
)
@@ -3876,7 +3896,7 @@ class RotaBuilder(object):
continue
if day in constraint_shift.days:
# Only apply if worker can work this shift
if worker.site not in constraint_shift.sites:
if not self.is_worker_eligible_for_shift(worker, constraint_shift):
continue
try:
works = self.model.works[
@@ -3891,7 +3911,7 @@ class RotaBuilder(object):
post_date = None
if post_date is not None and post_date in getattr(constraint, "exclude_dates", []):
continue
self.model.constraints.add(
1
>= works
@@ -3905,7 +3925,7 @@ class RotaBuilder(object):
)
if shiftname not in ignore_shifts
for w in workers
if w.site in constraint_shift.sites # Only workers who can work this shift
if self.is_worker_eligible_for_shift(w, constraint_shift) # Only workers who can work this shift
)
)
@@ -4181,7 +4201,7 @@ class RotaBuilder(object):
and shift.name in worker.assign_as_block_preferences
and (worker.id, week, shift.name)
in self.model.blocks_worker_shift_assigned
and worker.site in shift.sites
and self.is_worker_eligible_for_shift(worker, shift)
)
# Quadratic :(
@@ -4277,6 +4297,11 @@ class RotaBuilder(object):
"Invalid exact shift site",
f"Worker {worker.name} requested exact shifts for shift {s_name} but their site {worker.site} is not in the shift sites {s.sites}",
)
elif not self.is_worker_eligible_for_shift(worker, s):
self.add_warning(
"Invalid exact shift group",
f"Worker {worker.name} requested exact shifts for shift {s_name} but is not eligible for it",
)
except KeyError:
self.add_warning(
"Invalid exact shift",
@@ -4294,7 +4319,13 @@ class RotaBuilder(object):
self.add_warning("Worker/duplicate name", message)
self.workers_name_map[worker.name] = worker
if worker.site not in self.sites:
worker_has_valid_shift = False
for shift in self.shifts:
if self.is_worker_eligible_for_shift(worker, shift):
worker_has_valid_shift = True
break
if not worker_has_valid_shift:
message = f"Worker with name '{worker.name}' ({worker.id}) has no valid shifts (site: {worker.site})"
logger.warning(message)
self.add_warning("Worker/no valid shifts", message)
@@ -4873,8 +4904,7 @@ class RotaBuilder(object):
def get_shifts_for_worker(self, worker_id):
worker = self.get_worker_by_id(worker_id)
shifts = [shift for shift in self.shifts if worker.site in shift.sites]
return shifts
return [shift for shift in self.shifts if self.is_worker_eligible_for_shift(worker, shift)]
def get_shifts_with_constraint(self, constraint) -> List[SingleShift]:
return [shift for shift in self.shifts if shift.has_constraint(constraint)]
@@ -5066,8 +5096,46 @@ class RotaBuilder(object):
return group_workers
def is_worker_eligible_for_shift(self, worker: Worker, shift: SingleShift) -> bool:
"""Returns True if the worker is eligible for the shift based on site and group options."""
# 1. Group checks
worker_groups = getattr(worker, "groups", set())
if not isinstance(worker_groups, (set, list, tuple)):
worker_groups = {worker_groups} if worker_groups else set()
worker_groups_set = set(worker_groups)
# exclude_groups check: worker MUST NOT belong to any group in exclude_groups
if getattr(shift, "exclude_groups", None):
exclude_set = set(shift.exclude_groups)
if worker_groups_set.intersection(exclude_set):
return False
# Additive Site or Include Group check
has_site = worker.site in shift.sites
has_include_group = False
if getattr(shift, "include_groups", None):
include_set = set(shift.include_groups)
if worker_groups_set.intersection(include_set):
has_include_group = True
if getattr(shift, "include_groups", None):
if not (has_site or has_include_group):
return False
else:
if not has_site:
return False
return True
def get_workers_in_group(self, group_name: str) -> List[Worker]:
"""Returns a list of all workers belonging to the specified group."""
return [
w for w in self.workers
if group_name in getattr(w, "groups", set())
]
def get_workers_for_shift(self, shift: SingleShift) -> List[Worker]:
return [worker for worker in self.workers if worker.site in shift.sites]
return [worker for worker in self.workers if self.is_worker_eligible_for_shift(worker, shift)]
def get_workers_total_fte(self) -> float:
"""Does not take into account shift adjusted ftes"""
+1
View File
@@ -214,6 +214,7 @@ class Worker(BaseModel):
shift_fte_overrides: dict[str, int] = {} # Need checks to ensure shifts exist
exact_shifts: dict[str, int] = {} # Map shift_name to exact number of shifts
groups: set[str] = set() # Groups that the worker belongs to
weekend_shift_target_number: float = 0.0
avoid_shifts_on_dates: list[AvoidShiftOnDates] = []