# Import PuLP for linear programming (pulp), h3, Openrouteservice (ors)

# Import geographical data processing as necessary (Geopandas, Shapely, Geopy, Pyproj, ...)

# Declare necessary variables including resolution (in this case, H3 resolution 6 or 7), n_ped (number of hospitals aimed for)

# Calculate travel time from each H3 centroid to every hospital

def get_all_shortest_time(client, start_point, destinations):
    '''
    Computes shortest travel time from start point to multiple locations and
    returns the shortest travel time for all the given locations.

    Parameters
    ----------
    client : Client
        ORS client.
    start_point : tupel
        Start point.
    destinations : list
        List of tupels with coordinates.
    top : int
        Compute only distance to top n destinations with smallest air distance

    Returns
    -------
    Returns shortest travel time in min for given start point and all locations.
    '''
    locations = [start_point] + destinations

    distance_matrix = client.distance_matrix(locations, 
                                      profile='driving-car', 
                                      sources=[0],
                                      metrics=['duration', 'distance'], 
                                      units='m') 
    return distance_matrix

# Exclude centroids not in the area of interest and ones that are located on islands or water
# e.g., you may use spatial join in Geopandas excluding regions that are outside a given polygon surrounding a county's mainland

# Create Geopandas GeoDataFrame including hexagons in the area of interest on resolutions of choice
# in the below example, gdf has been declared as a GeoDataFrame with all H3 hexagons included in the analyses
# (one row each hexagon); merge gdf with population density of each H3 hexagon (density as a column)

# Read file with hospital locations. Create variable destination including the hospital coordinates (list of tupels with coordinates)

# Local Openrouteservice client
client = ors.Client(base_url = 'http://localhost:8080/ors', key='',
                    timeout=None,
                    retry_timeout=60) 

# To collect shortest travel time to next hospital for hex centroids

# in H3 version 4.x, the function h3_to_geo was renamed cell_to_latlng

distance_time_info = {}

for index, row in gdf.iterrows():

    shortest_time = get_all_shortest_time(client = client, 
                                          start_point = (h3.h3_to_geo(row["hexagon_id"])[1], h3.h3_to_geo(row["hexagon_id"])[0]),
                                          destinations = destinations)
    if type(shortest_time) is list:

        distance_time_info[row['hexagon_id']] = {"times": list(map(lambda x: round(x, ndigits=2), shortest_time)),
                                                 "density_times":  list(map(lambda x: round(x*row["density"], ndigits=2), shortest_time)),
                                                 "density": row["density"]
            }

# Extract the density-weighted travel times
costs_raw = [elem['density_times'] for key, elem in distance_time_info.items()]

# To use raw travel times instead of density-weighted times:
# costs_raw = [elem['times'] for key, elem in distance_time_info.items()]

# Define the optimization problem as a minimization problem
prob = pulp.LpProblem("HospitalAssignment", pulp.LpMinimize)

# Number of patients and hospitals derived from the distance matrix dimensions
num_patients = len(costs_raw)
num_hospitals = len(costs_raw[0])

# Define binary decision variables:
# x[i][j]: 1 if patient i is assigned to hospital j, 0 otherwise
x = pulp.LpVariable.dicts("x", (range(num_patients), range(num_hospitals)), 0, 1, pulp.LpBinary)

# y[j]: 1 if hospital j is open, 0 otherwise
y = pulp.LpVariable.dicts("y", range(num_hospitals), 0, 1, pulp.LpBinary)

# Objective function: Minimize the total weighted travel cost
prob += pulp.lpSum(costs_raw[i][j] * x[i][j] for i in range(num_patients) for j in range(num_hospitals))

# Constraint (1): Each patient is assigned to exactly one hospital
for i in range(num_patients):
    prob += pulp.lpSum(x[i][j] for j in range(num_hospitals)) == 1

# Constraint (2): Patients can only be assigned to open hospitals
for i in range(num_patients):
    for j in range(num_hospitals):
        prob += x[i][j] <= y[j]

# Constraint (3): The total number of open hospitals must not exceed the given limit
prob += pulp.lpSum(y[j] for j in range(num_hospitals)) <= n_ped

'''
# Scenario-specific constraints: for the model prioritizing pediatric hospitals
# Mark pediatric hospital indices (variable fixed_hospitals)
if n_ped > 315:
    prob += pulp.lpSum(y[j] for j in range(num_hospitals)) <= n_ped
    for j in fixed_hospitals:
        prob += y[j] == 1

# First filter the cost matrix W, excluding all non-pediatric hospitals
# Then the following condition may be used
if n_ped < 315:
    prob += pulp.lpSum(y[j] for j in range(num_hospitals)) == n_ped
'''

# Use the default Pulp CBC solver (choose function parameters as necessary)
solver = pulp.PULP_CBC_CMD()

# Solve the optimization problem
prob.solve(solver)

# Extract the indices of hospitals that are selected to remain open
hospitals_indices = [j for j in range(num_hospitals) if pulp.value(y[j]) == 1]

# Calculate travel time from each centroid to the corrersponding selected hospital using openrouteservice