refactor and add some more tests
This commit is contained in:
+49
-4
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user