refactor and add some more tests

This commit is contained in:
Ross
2022-05-16 19:35:38 +01:00
parent ff03358119
commit f887257a55
6 changed files with 427 additions and 56 deletions
+49 -4
View File
@@ -267,12 +267,13 @@ class RotaBuilder(object):
if not results.solver.status:
sys.exit(0)
def build_and_solve(self, options={"ratio": 0.1, "seconds": 1000, "threads": 10}):
def build_and_solve(self, options={"ratio": 0.1, "seconds": 1000, "threads": 10}, solve=True):
self.build_shifts()
self.build_workers()
self.build_model()
self.solve_model(options=options)
if solve:
self.solve_model(options=options)
def build_model(self):
# Initialize model
@@ -1916,7 +1917,7 @@ class RotaBuilder(object):
self.workers.append(worker)
def add_workers(self, workers: List) -> None:
"""Add multiple worker to the rota
"""Add multiple workers to the rota
Args:
workers (List(Worker)):
@@ -1932,6 +1933,21 @@ class RotaBuilder(object):
if not self.workers:
raise NoWorkers("Workers must be added prior to calling build_workers")
self.workers_id_map = {}
self.workers_name_map = {}
for worker in self.workers:
wid = worker.id
if wid in self.workers_id_map:
raise ValueError(f"Worker with id '{wid}' has been added twice")
self.workers_id_map[wid] = worker
if worker.name in self.workers_name_map:
raise ValueError(f"Worker with name '{worker.name}' has been added twice")
self.workers_name_map[worker.name] = worker
self.workers = sorted(self.workers)
self.full_time_equivalent = sum(w.fte_adj for w in self.workers)
@@ -2216,6 +2232,9 @@ class RotaBuilder(object):
s.extend(self.shifts_to_force_as_blocks())
return s
def get_worker_by_name(self, name: str) -> Worker:
return self.workers_name_map[name]
def get_workers_by_group(self) -> Dict[str, Worker]:
group_workers = defaultdict(set)
for worker in self.workers:
@@ -2315,6 +2334,7 @@ class RotaBuilder(object):
week_table[week][day][shift].append(worker.get_details())
return week_table
def get_worker_timetable(self):
works = self.model.works
timetable = {
@@ -2511,7 +2531,9 @@ class RotaBuilder(object):
worker_td = """<td title='Site: {site}' class='worker {site}'
data-nwds='{nwds}' data-site='{site}' data-worker='{name}'
data-fte='{fte}' data-fte_adj='{fte_adj}' data-end_date='{end_date}'
data-fte='{fte}' data-fte_adj='{fte_adj}'
data-start_date='{start_date}'
data-end_date='{end_date}'
data-worker-targets='{targets}' data-shift-counts='{worker_shift_counts}'
data-weekend-target='{weekend_target}'
data-remote-site='{remote_site}'
@@ -2527,6 +2549,7 @@ class RotaBuilder(object):
name=worker.name,
fte=worker.fte,
fte_adj=worker.fte_adj,
start_date=worker.start_date,
end_date=worker.end_date,
targets=worker_targets,
worker_shift_counts=json.dumps(shift_count_dict),
@@ -2645,6 +2668,22 @@ class RotaBuilder(object):
return timetable
def get_worker_shifts_by_date(self, worker: Worker) -> Dict:
shifts = {}
n = 0
for week, day in self.get_week_day_combinations():
d = self.start_date + datetime.timedelta(n)
n = n + 1
temp = ""
for shift in self.get_shift_names_by_week_day(week, day):
if self.model.works[worker.id, week, day, shift].value > 0:
shifts[d]=shift
return shifts
def get_worker_shift_list(self, worker: Worker) -> List:
shifts = []
@@ -2660,6 +2699,12 @@ class RotaBuilder(object):
return shifts
def get_worker_shift_list_string(self, worker: Worker) -> str:
shifts = self.get_worker_shift_list(worker)
# Convert shift to a string representation
return "".join([i if i != "" else "-" for i in shifts])
def get_shift_summary(self):
works = self.model.works
timetable = {