from pulp import *
from app.main.util.json_objects.games.stratedge.stratedge_result_dto import StratEdgeResultDTO
from app.main.util.json_objects.games.stratedge.market_result_dto import MarketResultDTO
import app.main.util.helpers as helpers


class StratEdgeSimulator:

    def __init__(self, round, stratedge_configuration, round_scenario_data,
                 previous_round_results):
        self.round = round
        self.stratedge_configuration = stratedge_configuration
        self.round_scenario_data = round_scenario_data
        self.previous_round_results = previous_round_results
        self.strategic_decisions_dict = None
        self.market_indices = None
        self.market_ids = None
        self.market_names = None
        self.market_demands = None
        self.competitor_indices = None
        self.competitor_names = None
        self.competitor_capacities = None
        self.competitor_costs = None
        self.competitor_capacities_before = None
        self.competitor_costs_before = None
        self.competitor_reserves = None
        self.competitor_total_investments = None
        self.competitor_total_volumes = None
        self.competitor_ebitda = None
        self.transportation_costs = None
        self.transportation_costs_before = None
        self.from_fixed_name_to_team_id = None
        self.prob = None
        self.variable_dict = None
        self.variables_by_competitor = None
        self.variables_by_market = None
        self.transformed_names_dict = None
        self.constraint_dict = None
        self.prices = None
        self.margins = None
        self.sales = None
        self.comp_volumes = None
        self.volumes = None
        self.market_margins = None
        self.round_index = int(round.name[-1])

    def compute(self):
        self.init_configuration()
        self.init_market_names_and_demands()
        self.init_competitors()
        self.apply_new_decisions()
        self.launch_price_war()
        self.construct_outputs()

    def init_configuration(self):
        self.strategic_decisions_dict = {}
        strategic_decisions = self.stratedge_configuration.strategic_decisions
        for strategic_decision in strategic_decisions:
            self.strategic_decisions_dict[strategic_decision.id] = strategic_decision

    def init_market_names_and_demands(self):
        self.market_indices = {}
        self.market_ids = {}
        self.market_names = ['', '', '', '']
        self.market_demands = [0, 0, 0, 0]
        for market in self.stratedge_configuration.markets:
            market_name = market.fixed_name
            market_index = int(market_name[-1]) - 1
            self.market_names[market_index] = market_name
            self.market_indices[market_name] = market_index
            self.market_ids[market_name] = market.id
            self.market_demands[market_index] = getattr(market, 'demand'+str(self.round_index))

    def init_competitors(self):
        self.competitor_indices = {}
        self.competitor_names = ['', '', '', '']
        self.competitor_capacities = [0,0,0,0]
        self.competitor_costs = [0,0,0,0]
        self.competitor_capacities_before = [0,0,0,0]
        self.competitor_costs_before = [0,0,0,0]
        self.competitor_reserves = [0,0,0,0]
        self.competitor_total_investments = [0,0,0,0]
        self.competitor_total_volumes = [0,0,0,0]
        self.competitor_ebitda = [0,0,0,0]
        self.transportation_costs = [[0,0],[0,0],[0,0],[0,0]]
        self.transportation_costs_before = [[0,0],[0,0],[0,0],[0,0]]
        if self.previous_round_results is not None and len(self.previous_round_results) > 0:
            for stratedge_result in self.previous_round_results:
                competitor_name = stratedge_result.se_team.competitor.fixed_name
                competitor_index = int(competitor_name[-1]) -1
                self.competitor_indices[competitor_name] = competitor_index
                self.competitor_names[competitor_index] = competitor_name
                self.competitor_capacities[competitor_index] = stratedge_result.capacity
                self.competitor_costs[competitor_index] = stratedge_result.cost
                self.competitor_capacities_before[competitor_index] = stratedge_result.capacity
                self.competitor_costs_before[competitor_index] = stratedge_result.cost
                self.competitor_reserves[competitor_index] = stratedge_result.reserve
                self.competitor_total_volumes[competitor_index] = stratedge_result.total_volume
                self.competitor_ebitda[competitor_index] = stratedge_result.ebitda
                self.competitor_total_investments[competitor_index] = 0
                for market_result in stratedge_result.market_results:
                    market_index = self.market_indices[market_result.se_market.fixed_name]
                    self.transportation_costs[competitor_index][market_index] = market_result.fret
                    self.transportation_costs_before[competitor_index][market_index] = market_result.fret
        else:
            for competitor in self.stratedge_configuration.competitors:
                competitor_name = competitor.fixed_name
                competitor_index = int(competitor_name[-1]) -1
                self.competitor_indices[competitor_name] = competitor_index
                self.competitor_names[competitor_index] = competitor_name
                self.competitor_capacities[competitor_index] = competitor.production_capacity
                self.competitor_costs[competitor_index] = competitor.production_cost
                self.competitor_capacities_before[competitor_index] = competitor.production_capacity
                self.competitor_costs_before[competitor_index] = competitor.production_cost
                self.competitor_reserves[competitor_index] = self.stratedge_configuration.initial_budget
                self.competitor_total_volumes[competitor_index] = 0
                self.competitor_ebitda[competitor_index] = 0
                self.competitor_total_investments[competitor_index] = 0
                self.transportation_costs[competitor_index][0] = competitor.fret_market1
                self.transportation_costs_before[competitor_index][0] = competitor.fret_market1
                self.transportation_costs[competitor_index][1] = competitor.fret_market2
                self.transportation_costs_before[competitor_index][1] = competitor.fret_market2

    def apply_new_decisions(self):
        self.from_fixed_name_to_team_id = {}
        team_scenarios_data = self.round_scenario_data['team_scenarios']
        for team_scenario_data in team_scenarios_data:
            competitor_name = team_scenario_data['competitor_fixed_name']
            team_id = team_scenario_data['team_id']
            self.from_fixed_name_to_team_id[competitor_name] = team_id
            competitor_index = self.competitor_indices[competitor_name]
            strategic_decisions_data = team_scenario_data['strategic_decisions']
            for strategic_decision_data in strategic_decisions_data:
                strategic_decision_id = strategic_decision_data['id']
                strategic_decision = self.strategic_decisions_dict[strategic_decision_id]
                if self.competitor_reserves[competitor_index] < strategic_decision.price:
                    continue
                self.competitor_costs[competitor_index] = round(
                    self.competitor_costs[competitor_index] * (1+strategic_decision.cost_impact))
                self.competitor_capacities[competitor_index] = round(
                    self.competitor_capacities[competitor_index] * (1+strategic_decision.capacity_impact))
                self.competitor_reserves[competitor_index] = round(
                    self.competitor_reserves[competitor_index] - strategic_decision.price)
                self.competitor_total_investments[competitor_index] = round(
                    self.competitor_total_investments[competitor_index] + strategic_decision.price)
                self.transportation_costs[competitor_index][0] = round(
                    self.transportation_costs[competitor_index][0] *(1+strategic_decision.fret_impact_market1))
                self.transportation_costs[competitor_index][1] = round(
                    self.transportation_costs[competitor_index][1] * (1 + strategic_decision.fret_impact_market2))

    def launch_price_war(self):
        self.compute_equilibrium()
        self.compute_volumes()
        self.compute_prices()
        self.compute_net_backs()

    def compute_equilibrium(self):
        self.init_problem_and_variables()
        self.construct_capacity_constraints()
        self.construct_demand_constraints()
        self.post_processing_solving()

    def init_problem_and_variables(self):
        self.prob = LpProblem("price_war", LpMinimize)
        self.prob.tmpDir = '/tmp'
        self.variable_dict = {}
        self.variables_by_competitor = {}
        self.variables_by_market = {}
        self.transformed_names_dict = {}
        obj_func = None
        for competitor_idx in range(4):
            competitor_name = self.competitor_names[competitor_idx]
            for market_idx in range(2):
                market_name = self.market_names[market_idx]
                variable_name = helpers.get_variable_name(competitor_name, market_name)
                if variable_name not in self.variable_dict:
                    x = LpVariable(variable_name, 0, min(self.competitor_capacities[competitor_idx], self.market_demands[market_idx]))
                    self.variable_dict[variable_name] = x
                    if competitor_name not in self.variables_by_competitor:
                        self.variables_by_competitor[competitor_name] = []
                    if market_name not in self.variables_by_market:
                        self.variables_by_market[market_name] = []
                    self.variables_by_competitor[competitor_name].append(x)
                    self.variables_by_market[market_name].append(x)
                    obj_func += (self.competitor_costs[competitor_idx] + self.transportation_costs[competitor_idx][market_idx]) * x
                    self.transformed_names_dict[variable_name] = x.name
        self.prob += obj_func

    def construct_capacity_constraints(self):
        self.constraint_dict = {}
        for competitor_name in self.variables_by_competitor:
            competitor_vars = self.variables_by_competitor[competitor_name]
            constraint_name = helpers.get_capacity_constraint_name(competitor_name)
            competitor_idx = self.competitor_indices[competitor_name]
            competitor_capacity_constraint = lpSum(competitor_vars) <= self.competitor_capacities[competitor_idx], constraint_name
            self.constraint_dict[constraint_name] = competitor_capacity_constraint
            self.prob += competitor_capacity_constraint
            self.transformed_names_dict[constraint_name] = competitor_capacity_constraint[0].name

    def construct_demand_constraints(self):
        for market_name in self.variables_by_market:
            market_vars = self.variables_by_market[market_name]
            constraint_name = helpers.get_demand_constraint_name(market_name)
            market_idx = self.market_indices[market_name]
            market_demand_constraint = lpSum(market_vars) >= self.market_demands[market_idx], constraint_name
            self.constraint_dict[constraint_name] = market_demand_constraint
            self.prob += market_demand_constraint
            self.transformed_names_dict[constraint_name] = market_demand_constraint[0].name

    def post_processing_solving(self):
        self.prob.solve()

    def compute_volumes(self):
        self.comp_volumes = [0, 0, 0, 0]
        self.volumes = [[0, 0], [0, 0], [0, 0], [0, 0]]
        for competitor_idx in range(4):
            competitor_name = self.competitor_names[competitor_idx]
            for market_idx in range(2):
                market_name = self.market_names[market_idx]
                variable_name = helpers.get_variable_name(competitor_name, market_name)
                variable = self.variable_dict[variable_name]
                volume = min(self.competitor_capacities[competitor_idx] - self.comp_volumes[competitor_idx], variable.varValue)
                self.comp_volumes[competitor_idx] += volume
                self.volumes[competitor_idx][market_idx] = volume

    def compute_prices(self):
        self.prices = [0,0]
        for market_idx in range(2):
            market_name = self.market_names[market_idx]
            constraint_name = helpers.get_demand_constraint_name(market_name)
            transformed_constraint_name = self.transformed_names_dict[constraint_name]
            market_demand_constraint = self.prob.constraints[transformed_constraint_name]
            self.prices[market_idx] = market_demand_constraint.pi + 5

    def compute_net_backs(self):
        self.sales = [0, 0, 0, 0]
        self.margins = [0, 0, 0, 0]
        self.market_margins = [[0, 0], [0, 0], [0, 0],[0, 0]]
        for competitor_idx in range(4):
            for market_idx in range(2):
                volume = self.volumes[competitor_idx][market_idx]
                self.sales[competitor_idx] += volume * self.prices[market_idx]
                market_margin = volume * (self.prices[market_idx] - (self.competitor_costs[competitor_idx] +
                                                           self.transportation_costs[competitor_idx][market_idx]))
                self.market_margins[competitor_idx][market_idx] = market_margin
                self.margins[competitor_idx] += market_margin

    def construct_outputs(self):
        stratedge_result_dtos = []
        for competitor_idx in range(4):
            competitor_name = self.competitor_names[competitor_idx]
            stratedge_result_dto = self.construct_stratedge_result_dto(competitor_idx, competitor_name)
            stratedge_result_dtos.append(stratedge_result_dto)
        stratedge_results_data = [x.to_json() for x in stratedge_result_dtos]
        self.round_scenario_data['scenario_results'] = stratedge_results_data

    def construct_stratedge_result_dto(self, competitor_idx, competitor_name):
        stratedge_result_dto = StratEdgeResultDTO()
        stratedge_result_dto.team_id = self.from_fixed_name_to_team_id[competitor_name]
        stratedge_result_dto.total_investment = self.competitor_total_investments[competitor_idx]
        stratedge_result_dto.cost = self.competitor_costs[competitor_idx]
        stratedge_result_dto.cost_variation = round(((self.competitor_costs[competitor_idx]
                                                     / self.competitor_costs_before[competitor_idx])-1)*100.0)
        stratedge_result_dto.capacity = self.competitor_capacities[competitor_idx]
        stratedge_result_dto.capacity_variation = round(((self.competitor_capacities[competitor_idx]
                                                        / self.competitor_capacities_before[competitor_idx])-1)*100.0)
        stratedge_result_dto.reserve = round(self.competitor_reserves[competitor_idx] + self.margins[competitor_idx])
        stratedge_result_dto.ca = round(self.sales[competitor_idx])
        stratedge_result_dto.ebitda = round(self.margins[competitor_idx])
        stratedge_result_dto.ebitda_variation = 0
        if self.competitor_ebitda[competitor_idx] != 0:
            stratedge_result_dto.ebitda_variation = round(((self.margins[competitor_idx]
                                                           / self.competitor_ebitda[competitor_idx])-1)*100.0)
        stratedge_result_dto.total_volume = round(self.comp_volumes[competitor_idx])
        stratedge_result_dto.total_volume_variation = 0
        if self.competitor_total_volumes[competitor_idx] != 0:
            stratedge_result_dto.total_volume_variation = round(((self.comp_volumes[competitor_idx]
                                                            / self.competitor_total_volumes[competitor_idx])-1)*100.0)
        market_result_dtos = []
        for market_idx in range(2):
            market_name = self.market_names[market_idx]
            market_result_dto = self.construct_market_result_dto(competitor_idx, market_idx, market_name)
            market_result_dtos.append(market_result_dto)
        stratedge_result_dto.market_results = market_result_dtos
        return stratedge_result_dto

    def construct_market_result_dto(self, competitor_idx, market_idx, market_name):
        market_result_dto = MarketResultDTO()
        market_result_dto.market_id = self.market_ids[market_name]
        market_result_dto.price = self.prices[market_idx]
        market_result_dto.cost = round(self.competitor_costs[competitor_idx] + self.transportation_costs[competitor_idx][market_idx])
        market_result_dto.fret = round(self.transportation_costs[competitor_idx][market_idx])

        market_result_dto.fret_variation = round((((self.competitor_costs[competitor_idx] + self.transportation_costs[competitor_idx][market_idx])
                                            / (self.competitor_costs_before[competitor_idx] + self.transportation_costs_before[competitor_idx][market_idx]))-1) * 100.0)
        market_result_dto.volume = round(self.volumes[competitor_idx][market_idx])
        market_result_dto.margin = round(self.market_margins[competitor_idx][market_idx])
        return market_result_dto
