from __future__ import annotations

import math
import os
import platform
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Dict, Iterable, List, Sequence, Tuple
from concurrent.futures import ProcessPoolExecutor, as_completed

import numpy as np
import pandas as pd

OUT = Path(__file__).resolve().parents[1] / 'outputs'
OUT.mkdir(parents=True, exist_ok=True)

@dataclass(frozen=True)
class Job:
    jid: int
    release: float
    pickup: Tuple[int, int]
    dropoff: Tuple[int, int]
    due: float
    priority: int
    travel_multiplier: float
    handling_time: float
    energy_multiplier: float

@dataclass
class Vehicle:
    vid: int
    node: Tuple[int, int]
    available: float
    soc: float
    busy_time: float = 0.0
    empty_dist: float = 0.0
    loaded_dist: float = 0.0
    charge_service: float = 0.0
    charge_queue: float = 0.0
    jobs: int = 0

PROFILES = {
    'nominal': dict(arrival_rate=.34, block_prob=.015, block_count=1, speed_factor=1.00, due_factor=2.40),
    'demand_surge': dict(arrival_rate=.46, block_prob=.020, block_count=1, speed_factor=1.04, due_factor=2.15),
    'aisle_disruption': dict(arrival_rate=.35, block_prob=.120, block_count=3, speed_factor=1.10, due_factor=2.30),
    'combined_stress': dict(arrival_rate=.45, block_prob=.100, block_count=3, speed_factor=1.14, due_factor=2.05),
}
PROFILE_INDEX = {p: i for i, p in enumerate(PROFILES)}

W, H = 7, 6
STATIONS = [(0,0),(0,H-1),(W-1,0),(W-1,H-1),(W//2,0),(W//2,H-1),(1,2),(W-2,3),(2,H-2),(W-3,1)]
ALL_CHARGERS = [(0,0),(W-1,H-1)]
PARKING = [(0,H-1),(W-1,0),(W//2,0),(W//2,H-1)]
BUCKET_LENGTH = 15.0
MAX_BUCKETS = 260
CANDIDATE_LIMIT = 18
ENERGY_PER_DISTANCE = 1.15
SAFETY_SOC = 8.0
CHARGE_RATE = 2.2

BASE_WEIGHTS = {
    'empty': 1.00,
    'inflation': 1.25,
    'reserve': 2.50,
    'tardiness': 2.00,
    'charger_wait': 0.50,
    'priority': 0.80,
}
SCALES = {
    'empty': 10.0,
    'inflation': 5.0,
    'reserve': 20.0,
    'tardiness': 50.0,
    'charger_wait': 50.0,
    'priority': 2.0,
}

MAIN_POLICIES = [
    'FIFO-nearest',
    'EDD-nearest',
    'Energy-core',
    'RCRD-no-route',
    'RCRD-no-wait',
    'Full RCRD',
]


def stable_rng(seed: int, profile: str, stream: int) -> np.random.Generator:
    return np.random.default_rng(np.random.SeedSequence([seed, PROFILE_INDEX[profile], stream]))


def base_dist(a: Tuple[int,int], b: Tuple[int,int]) -> float:
    return abs(a[0]-b[0]) + abs(a[1]-b[1])


def dist(a: Tuple[int,int], b: Tuple[int,int], zones: Sequence[Tuple[str,int,float]] = ()) -> float:
    d = float(base_dist(a,b))
    for axis, k, penalty in zones:
        if axis == 'x' and min(a[0],b[0]) < k <= max(a[0],b[0]):
            d += penalty
        if axis == 'y' and min(a[1],b[1]) < k <= max(a[1],b[1]):
            d += penalty
    return d


def make_chargers(n_chargers: int) -> Tuple[Tuple[int,int], ...]:
    if n_chargers not in (1,2):
        raise ValueError('n_chargers must be 1 or 2')
    return tuple(ALL_CHARGERS[:n_chargers])


def nearest_charger(node: Tuple[int,int], chargers: Sequence[Tuple[int,int]], zones=()):
    return min(((dist(node,c,zones), c) for c in chargers), key=lambda x: x[0])


def generate_jobs(seed: int, profile: str, n_jobs: int) -> List[Job]:
    rng = stable_rng(seed, profile, 1)
    p = PROFILES[profile]
    t = 0.0
    jobs: List[Job] = []
    for j in range(n_jobs):
        t += rng.exponential(1 / p['arrival_rate'])
        ii = rng.choice(len(STATIONS), 2, replace=False)
        a, b = STATIONS[int(ii[0])], STATIONS[int(ii[1])]
        loaded = base_dist(a,b)
        r = rng.random()
        priority = 2 if r < .18 else (1 if r < .48 else 0)
        allowance = p['due_factor'] * (loaded + 4) * (0.82 if priority == 2 else (0.92 if priority == 1 else 1.0)) + rng.uniform(4,12)
        travel_multiplier = max(.75, rng.normal(p['speed_factor'], .06))
        handling = max(.8, rng.normal(2.0, .35))
        energy_multiplier = ENERGY_PER_DISTANCE * max(.88, 1 + rng.normal(0, .035))
        jobs.append(Job(j, t, a, b, t+allowance, priority, travel_multiplier, handling, energy_multiplier))
    return jobs


def generate_disruption_schedule(seed: int, profile: str) -> Tuple[Tuple[Tuple[str,int,float], ...], ...]:
    rng = stable_rng(seed, profile, 2)
    p = PROFILES[profile]
    schedule = []
    for bucket in range(MAX_BUCKETS):
        t = bucket * BUCKET_LENGTH
        circadian = 1 + .45 * math.sin((t/240.0) * math.pi)
        zones = []
        if rng.random() < min(.40, p['block_prob'] * circadian):
            count = p['block_count'] + int(rng.random() < .30)
            for _ in range(count):
                if rng.random() < .5:
                    zones.append(('x', int(rng.integers(1, W-1)), float(rng.uniform(1.5,4.5))))
                else:
                    zones.append(('y', int(rng.integers(1, H-1)), float(rng.uniform(1.5,4.5))))
        schedule.append(tuple(zones))
    return tuple(schedule)


def zones_at(schedule, t: float):
    idx = min(int(max(0.0,t) // BUCKET_LENGTH), len(schedule)-1)
    return schedule[idx]


def initial_soc(seed: int, profile: str, fleet: int) -> np.ndarray:
    rng = stable_rng(seed, profile, 3)
    return rng.uniform(72,100,size=fleet)


def estimate_features(v: Vehicle, j: Job, now: float, zones, chargers, charger_avail):
    empty_eff = dist(v.node, j.pickup, zones)
    loaded_eff = dist(j.pickup, j.dropoff, zones)
    empty_nom = base_dist(v.node, j.pickup)
    loaded_nom = base_dist(j.pickup, j.dropoff)
    total_eff = empty_eff + loaded_eff
    inflation = total_eff - (empty_nom + loaded_nom)
    post_soc = v.soc - total_eff * ENERGY_PER_DISTANCE
    dc, charger = nearest_charger(j.dropoff, chargers, zones)
    reserve_required = SAFETY_SOC + dc * ENERGY_PER_DISTANCE
    reserve_headroom = post_soc - reserve_required
    reserve_soft_penalty = max(0.0, 12.0 - reserve_headroom)
    expected_finish = now + total_eff * PROFILES_CURRENT['speed_factor'] + 2.0
    tardiness = max(0.0, expected_finish - j.due)
    charger_wait = max(0.0, charger_avail[charger] - expected_finish)
    feasible = reserve_headroom >= -1e-9
    return {
        'empty': empty_eff,
        'inflation': inflation,
        'reserve': reserve_soft_penalty,
        'tardiness': tardiness,
        'charger_wait': charger_wait,
        'priority': float(j.priority),
        'feasible': feasible,
        'total_eff': total_eff,
        'loaded_eff': loaded_eff,
        'reserve_headroom': reserve_headroom,
    }

# Worker-local profile settings. It avoids repeatedly passing a dict into a hot inner loop.
PROFILES_CURRENT = PROFILES['nominal']


def weighted_score(policy: str, features: Dict[str,float], weights: Dict[str,float]):
    active = {
        'Energy-core': {'empty','reserve','tardiness','priority'},
        'RCRD-no-route': {'empty','reserve','tardiness','charger_wait','priority'},
        'RCRD-no-wait': {'empty','inflation','reserve','tardiness','priority'},
        'Full RCRD': {'empty','inflation','reserve','tardiness','charger_wait','priority'},
    }[policy]
    value = 0.0
    for name in ('empty','inflation','reserve','tardiness','charger_wait'):
        if name in active:
            value += weights[name] * features[name] / SCALES[name]
    if 'priority' in active:
        value -= weights['priority'] * features['priority'] / SCALES['priority']
    return value


def adaptive_target(backlog: int, fixed: bool = False) -> float:
    if fixed:
        return 85.0
    return min(90.0, 80.0 + 0.8 * backlog)


def choose_vehicle_to_charge(free: Sequence[Vehicle], now: float, zones, chargers, charger_avail, backlog: int, fixed_target: bool):
    target = adaptive_target(backlog, fixed_target)
    candidates = []
    for v in free:
        d, c = nearest_charger(v.node, chargers, zones)
        arrive = now + d
        start = max(arrive, charger_avail[c])
        soc_arrival = max(0.0, v.soc - d * ENERGY_PER_DISTANCE)
        finish = start + max(0.0, target - soc_arrival) / CHARGE_RATE
        candidates.append((finish, -v.soc, v.vid, v))
    return min(candidates)[-1]


def charge(v: Vehicle, now: float, zones, chargers, charger_avail, target: float):
    d, c = nearest_charger(v.node, chargers, zones)
    arrive = now + d
    start = max(arrive, charger_avail[c])
    soc_arrival = max(0.0, v.soc - d * ENERGY_PER_DISTANCE)
    service = max(0.0, target - soc_arrival) / CHARGE_RATE
    finish = start + service
    charger_avail[c] = finish
    v.empty_dist += d
    v.busy_time += finish - now
    v.charge_service += service
    v.charge_queue += max(0.0, start-arrive)
    v.soc = target
    v.node = c
    v.available = finish


def run_one(seed: int, profile: str, policy: str, n_jobs: int = 450, fleet: int = 8, n_chargers: int = 2,
            weights: Dict[str,float] | None = None, candidate_limit: int = CANDIDATE_LIMIT,
            fixed_target: bool = False, config_label: str = 'base') -> Dict[str,float]:
    global PROFILES_CURRENT
    PROFILES_CURRENT = PROFILES[profile]
    weights = dict(BASE_WEIGHTS if weights is None else weights)
    jobs = generate_jobs(seed, profile, n_jobs)
    schedule = generate_disruption_schedule(seed, profile)
    chargers = make_chargers(n_chargers)
    socs = initial_soc(seed, profile, fleet)
    vehicles = [Vehicle(i, PARKING[i % len(PARKING)], 0.0, float(socs[i])) for i in range(fleet)]
    charger_avail = {c: 0.0 for c in chargers}
    waiting: List[Job] = []
    idx = 0
    done = []
    now = 0.0
    decisions = 0
    decision_times_ns = []
    disruption_assignments = 0
    infeasible_assignments = 0

    while len(done) < n_jobs:
        next_release = jobs[idx].release if idx < n_jobs else float('inf')
        next_available = min(v.available for v in vehicles)
        free = [v for v in vehicles if v.available <= now + 1e-9]
        if not waiting:
            now = next_release
        elif not free:
            now = min(next_release, next_available)
        while idx < n_jobs and jobs[idx].release <= now + 1e-9:
            waiting.append(jobs[idx]); idx += 1
        free = [v for v in vehicles if v.available <= now + 1e-9]
        if not waiting or not free:
            if idx >= n_jobs and waiting:
                now = next_available
            continue

        zones = zones_at(schedule, now)
        # Conventional rules evaluate the full queue so that FIFO and EDD retain
        # their standard meanings. Score-based policies evaluate a bounded set
        # of the most urgent jobs, which reflects an implementable controller.
        if policy == 'FIFO-nearest':
            cand = sorted(waiting, key=lambda j: (j.release, -j.priority, j.due, j.jid))
        elif policy == 'EDD-nearest':
            cand = sorted(waiting, key=lambda j: (j.due, -j.priority, j.release, j.jid))
        else:
            cand = sorted(waiting, key=lambda j: (j.due, -j.priority, j.release, j.jid))[:candidate_limit]
        t0 = time.perf_counter_ns()
        feasible_pairs = []
        for v in free:
            for j in cand:
                f = estimate_features(v, j, now, zones, chargers, charger_avail)
                if not f['feasible']:
                    continue
                if policy == 'FIFO-nearest':
                    score = (j.release, f['empty'], -f['reserve_headroom'], v.vid, j.jid)
                elif policy == 'EDD-nearest':
                    score = (j.due, f['empty'], -f['reserve_headroom'], v.vid, j.jid)
                else:
                    score = (weighted_score(policy, f, weights), j.due, f['empty'], v.vid, j.jid)
                feasible_pairs.append((score, v, j, f))
        decisions += 1
        if not feasible_pairs:
            v = choose_vehicle_to_charge(free, now, zones, chargers, charger_avail, len(waiting), fixed_target)
            target = adaptive_target(len(waiting), fixed_target)
            charge(v, now, zones, chargers, charger_avail, target)
            decision_times_ns.append(time.perf_counter_ns() - t0)
            continue
        _, v, j, f = min(feasible_pairs, key=lambda x: x[0])
        decision_times_ns.append(time.perf_counter_ns() - t0)
        if f['inflation'] > 1e-12:
            disruption_assignments += 1
        if not f['feasible']:
            infeasible_assignments += 1

        total = f['total_eff']
        empty = f['empty']
        loaded = f['loaded_eff']
        duration = total * j.travel_multiplier + j.handling_time
        finish = now + duration
        energy = total * j.energy_multiplier
        v.soc = max(0.0, v.soc-energy)
        v.node = j.dropoff
        v.available = finish
        v.busy_time += duration
        v.empty_dist += empty
        v.loaded_dist += loaded
        v.jobs += 1
        tard = max(0.0, finish-j.due)
        done.append((now-j.release, tard, tard <= 1e-9, finish-j.release, j.priority, empty, loaded, finish, f['inflation']))
        waiting.remove(j)

    df = pd.DataFrame(done, columns=['wait','tard','on_time','cycle','priority','empty','loaded','finish','inflation'])
    horizon = max(v.available for v in vehicles)
    empty = sum(v.empty_dist for v in vehicles)
    loaded = sum(v.loaded_dist for v in vehicles)
    dts = np.asarray(decision_times_ns, dtype=float) / 1000.0
    return dict(
        seed=seed, profile=profile, policy=policy, config=config_label, fleet=fleet, chargers=n_chargers,
        mean_tardiness=df.tard.mean(), p95_tardiness=df.tard.quantile(.95), on_time_rate=df.on_time.mean(),
        mean_wait=df.wait.mean(), p95_cycle=df.cycle.quantile(.95), empty_distance=empty, loaded_distance=loaded,
        empty_ratio=empty/(empty+loaded), utilization=sum(v.busy_time for v in vehicles)/(fleet*horizon),
        total_charge_service=sum(v.charge_service for v in vehicles), total_charge_queue=sum(v.charge_queue for v in vehicles),
        horizon=horizon, mean_route_inflation=df.inflation.mean(), disruption_assignment_share=disruption_assignments/n_jobs,
        infeasible_assignments=infeasible_assignments, decisions=decisions,
        dispatch_us_mean=dts.mean(), dispatch_us_p95=np.quantile(dts,.95), dispatch_us_max=dts.max(),
    )


def worker(kwargs):
    return run_one(**kwargs)


def parallel_run(tasks: List[Dict], label: str, max_workers: int | None = None) -> pd.DataFrame:
    rows = []
    max_workers = max_workers or max(1, min(os.cpu_count() or 2, 12))
    started = time.time()
    with ProcessPoolExecutor(max_workers=max_workers) as ex:
        futures = [ex.submit(worker, task) for task in tasks]
        for i, fut in enumerate(as_completed(futures), 1):
            rows.append(fut.result())
            if i % max(20, len(tasks)//10) == 0 or i == len(tasks):
                print(f'{label}: {i}/{len(tasks)} completed in {time.time()-started:.1f}s', flush=True)
    return pd.DataFrame(rows)


def bootstrap_paired(raw: pd.DataFrame, baseline: str, proposed: str, metrics: Sequence[str], nboot: int = 5000):
    rng = np.random.default_rng(20260731)
    p = raw[raw.policy == proposed].set_index(['seed','profile','config'])
    b = raw[raw.policy == baseline].set_index(['seed','profile','config'])
    rows = []
    for metric in metrics:
        if metric == 'on_time_rate':
            vals = (p[metric]-b[metric]).dropna().values
        else:
            vals = (b[metric]-p[metric]).dropna().values
        boots = np.array([rng.choice(vals, len(vals), replace=True).mean() for _ in range(nboot)])
        rows.append(dict(baseline=baseline, proposed=proposed, metric=metric, n=len(vals),
                         mean_improvement=vals.mean(), ci_low=np.quantile(boots,.025), ci_high=np.quantile(boots,.975)))
    return rows


def summarize(raw: pd.DataFrame, group_cols: Sequence[str]):
    metrics = ['mean_tardiness','p95_tardiness','on_time_rate','mean_wait','p95_cycle','empty_ratio','utilization',
               'total_charge_service','total_charge_queue','horizon','dispatch_us_mean','dispatch_us_p95']
    out = raw.groupby(list(group_cols))[metrics].agg(['mean','std']).reset_index()
    out.columns = ['_'.join(c).rstrip('_') for c in out.columns]
    return out


def run_main():
    tasks = []
    for seed in range(100,140):
        for profile in PROFILES:
            for policy in MAIN_POLICIES:
                tasks.append(dict(seed=seed, profile=profile, policy=policy, config_label='base'))
    raw = parallel_run(tasks, 'main')
    raw.to_csv(OUT/'main_results_by_run.csv', index=False)
    summarize(raw, ['profile','policy']).to_csv(OUT/'main_results_by_profile.csv', index=False)
    summarize(raw, ['policy']).to_csv(OUT/'main_results_overall.csv', index=False)
    comps = []
    for baseline in ['FIFO-nearest','EDD-nearest','Energy-core','RCRD-no-route','RCRD-no-wait']:
        comps.extend(bootstrap_paired(raw, baseline, 'Full RCRD',
                                      ['mean_tardiness','p95_tardiness','on_time_rate','empty_ratio','horizon']))
    pd.DataFrame(comps).to_csv(OUT/'main_paired_bootstrap.csv', index=False)
    return raw


def run_robustness():
    configs = [
        ('fleet_constrained', 6, 2),
        ('charger_constrained', 8, 1),
        ('capacity_expanded', 10, 2),
    ]
    tasks = []
    for seed in range(200,220):
        for profile in ['demand_surge','combined_stress']:
            for label, fleet, n_chargers in configs:
                for policy in ['Energy-core','Full RCRD']:
                    tasks.append(dict(seed=seed, profile=profile, policy=policy, fleet=fleet,
                                      n_chargers=n_chargers, config_label=label))
    raw = parallel_run(tasks, 'robustness')
    raw.to_csv(OUT/'robustness_results_by_run.csv', index=False)
    summarize(raw, ['config','profile','policy']).to_csv(OUT/'robustness_summary.csv', index=False)
    comps = bootstrap_paired(raw, 'Energy-core', 'Full RCRD',
                             ['mean_tardiness','p95_tardiness','on_time_rate','empty_ratio','horizon'])
    pd.DataFrame(comps).to_csv(OUT/'robustness_paired_bootstrap.csv', index=False)
    return raw


def run_sensitivity():
    variants = [('base','none',1.0,dict(BASE_WEIGHTS))]
    for name in BASE_WEIGHTS:
        for factor in (.75,1.25):
            w = dict(BASE_WEIGHTS); w[name] *= factor
            variants.append((f'{name}_{factor:.2f}',name,factor,w))
    tasks = []
    meta = {}
    for label,name,factor,w in variants:
        meta[label] = (name,factor,w)
        for seed in range(300,310):
            for profile in ['demand_surge','combined_stress']:
                tasks.append(dict(seed=seed, profile=profile, policy='Full RCRD', weights=w, config_label=label))
    raw = parallel_run(tasks, 'sensitivity')
    raw['weight_name'] = raw.config.map(lambda x: meta[x][0])
    raw['weight_factor'] = raw.config.map(lambda x: meta[x][1])
    raw.to_csv(OUT/'sensitivity_results_by_run.csv', index=False)
    summarize(raw, ['config','weight_name','weight_factor','profile']).to_csv(OUT/'sensitivity_summary.csv', index=False)
    return raw


def write_environment():
    text = [
        f'python={platform.python_version()}',
        f'platform={platform.platform()}',
        f'processor={platform.processor()}',
        f'cpu_count={os.cpu_count()}',
        f'numpy={np.__version__}',
        f'pandas={pd.__version__}',
        f'candidate_limit={CANDIDATE_LIMIT}',
        f'bucket_length={BUCKET_LENGTH}',
    ]
    (OUT/'environment.txt').write_text('\n'.join(text)+'\n')


def main():
    start = time.time()
    write_environment()
    run_main()
    run_robustness()
    run_sensitivity()
    print(f'All experiments completed in {time.time()-start:.1f}s', flush=True)

if __name__ == '__main__':
    main()
