feat: add group management for workers and shifts, including validation for overlapping groups
This commit is contained in:
+91
-23
@@ -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"""
|
||||
|
||||
@@ -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] = []
|
||||
|
||||
Reference in New Issue
Block a user