#!/usr/bin/env python3

# VRP_Flotten_Routing.py
"""
Kapitel Graphen: Capacitated Vehicle Routing Problem with Time Windows (CVRPTW)
mit der Routing-Bibliothek von Google OR-Tools.

Eigenschaften:
  * Eingabedaten werden vorab auf Plausibilitaet geprueft (Kapazitaet
    ausreichend? Zeitfenster erreichbar?)
  * Fahrzeit und Servicezeit werden getrennt ausgewiesen
  * Ausgabe als lesbarer Tourenplan mit Ankunftszeiten
  * Kennzahlen: Auslastung, Leerfahrten, Wartezeit
"""

import numpy as np
from ortools.constraint_solver import pywrapcp, routing_enums_pb2

SERVICEZEIT = 10          # Minuten je Kundenstopp
WARTEZEIT_MAX = 60        # zulaessige Wartezeit bei zu frueher Ankunft
SCHICHTLAENGE = 600       # Minuten


def erzeuge_daten(seed: int = 42):
    """Synthetische, aber reproduzierbare Instanz: 1 Depot + 16 Kunden."""
    anzahl_orte = 17
    rng = np.random.default_rng(seed)
    koordinaten = rng.random((anzahl_orte, 2)) * 100      # 100 x 100 km Raster

    distanz = np.zeros((anzahl_orte, anzahl_orte), dtype=int)
    for i in range(anzahl_orte):
        for j in range(anzahl_orte):
            distanz[i][j] = int(np.linalg.norm(koordinaten[i] - koordinaten[j]))

    return {
        "distanzmatrix": distanz.tolist(),
        "zeitfenster": [
            (0, SCHICHTLAENGE),                                   # 0: Depot
            (30, 120),  (60, 180),  (100, 240), (150, 300),       # Kunden 1-4
            (60, 180),  (120, 240), (200, 360), (300, 450),       # Kunden 5-8
            (180, 300), (240, 360), (300, 480), (360, 500),       # Kunden 9-12
            (60, 200),  (120, 300), (240, 400), (300, 550),       # Kunden 13-16
        ],
        "bedarfe": [0, 2, 3, 1, 4, 2, 2, 3, 1, 2, 4, 3, 2, 1, 2, 3, 2],
        "kapazitaeten": [10, 10, 10, 10],
        "anzahl_fahrzeuge": 4,
        "depot": 0,
    }


def pruefe_daten(daten) -> None:
    """Vorabdiagnose - fangt die haeufigsten Ursachen fuer 'keine Loesung' ab."""
    gesamtbedarf = sum(daten["bedarfe"])
    gesamtkapazitaet = sum(daten["kapazitaeten"])
    print(f"Gesamtbedarf {gesamtbedarf} Einheiten | "
          f"Flottenkapazitaet {gesamtkapazitaet} Einheiten | "
          f"Auslastung {gesamtbedarf / gesamtkapazitaet * 100:.0f} %")
    if gesamtbedarf > gesamtkapazitaet:
        raise SystemExit("UNLOESBAR: Der Bedarf uebersteigt die Flottenkapazitaet.")

    d = daten["distanzmatrix"]
    for kunde, (fruehestens, spaetestens) in enumerate(daten["zeitfenster"]):
        if kunde == 0:
            continue
        direktfahrt = d[0][kunde]
        if direktfahrt > spaetestens:
            raise SystemExit(
                f"UNLOESBAR: Kunde {kunde} ist erst nach {direktfahrt} min erreichbar, "
                f"sein Zeitfenster endet aber bei {spaetestens} min.")
    print("Vorabpruefung bestanden: Kapazitaet und Zeitfenster sind grundsaetzlich machbar.")


def loese_cvrptw(zeitlimit_s: int = 5):
    daten = erzeuge_daten()
    pruefe_daten(daten)

    manager = pywrapcp.RoutingIndexManager(
        len(daten["distanzmatrix"]), daten["anzahl_fahrzeuge"], daten["depot"])
    routing = pywrapcp.RoutingModel(manager)

    # --- Fahrzeit + Servicezeit als Kantengewicht ------------------------
    def zeit_callback(von_index, nach_index):
        von = manager.IndexToNode(von_index)
        nach = manager.IndexToNode(nach_index)
        service = SERVICEZEIT if von != daten["depot"] else 0
        return daten["distanzmatrix"][von][nach] + service

    zeit_index = routing.RegisterTransitCallback(zeit_callback)
    routing.SetArcCostEvaluatorOfAllVehicles(zeit_index)

    # --- Kapazitaetsdimension ---------------------------------------------
    def bedarf_callback(von_index):
        return daten["bedarfe"][manager.IndexToNode(von_index)]

    bedarf_index = routing.RegisterUnaryTransitCallback(bedarf_callback)
    routing.AddDimensionWithVehicleCapacity(
        bedarf_index, 0, daten["kapazitaeten"], True, "Kapazitaet")

    # --- Zeitdimension mit Zeitfenstern -----------------------------------
    routing.AddDimension(zeit_index, WARTEZEIT_MAX, SCHICHTLAENGE, False, "Zeit")
    zeit_dimension = routing.GetDimensionOrDie("Zeit")
    for ort, (fruehestens, spaetestens) in enumerate(daten["zeitfenster"]):
        zeit_dimension.CumulVar(manager.NodeToIndex(ort)).SetRange(fruehestens, spaetestens)

    # --- Suchparameter -----------------------------------------------------
    parameter = pywrapcp.DefaultRoutingSearchParameters()
    parameter.first_solution_strategy = (
        routing_enums_pb2.FirstSolutionStrategy.PATH_CHEAPEST_ARC)
    parameter.local_search_metaheuristic = (
        routing_enums_pb2.LocalSearchMetaheuristic.GUIDED_LOCAL_SEARCH)
    parameter.time_limit.seconds = zeitlimit_s

    loesung = routing.SolveWithParameters(parameter)
    if not loesung:
        print("Keine zulaessige Routenfuehrung gefunden.")
        return

    # --- Auswertung --------------------------------------------------------
    print("\n" + "=" * 84)
    print("         OPTIMIERTER TOURENPLAN (CVRPTW)")
    print("=" * 84)

    gesamtzeit = gesamtfracht = gesamtdistanz = 0
    kapazitaet = daten["kapazitaeten"]

    for fahrzeug in range(daten["anzahl_fahrzeuge"]):
        index = routing.Start(fahrzeug)
        if routing.IsEnd(loesung.Value(routing.NextVar(index))):
            print(f"\nFahrzeug {fahrzeug + 1}: nicht eingesetzt")
            continue

        stationen, fracht, distanz = [], 0, 0
        while not routing.IsEnd(index):
            knoten = manager.IndexToNode(index)
            ankunft = loesung.Min(zeit_dimension.CumulVar(index))
            fracht += daten["bedarfe"][knoten]
            bezeichnung = "Depot" if knoten == 0 else f"K{knoten}"
            stationen.append(f"{bezeichnung}@{ankunft}")
            naechster = loesung.Value(routing.NextVar(index))
            distanz += daten["distanzmatrix"][knoten][manager.IndexToNode(naechster)]
            index = naechster

        endzeit = loesung.Min(zeit_dimension.CumulVar(index))
        stationen.append(f"Depot@{endzeit}")
        gesamtzeit += endzeit
        gesamtfracht += fracht
        gesamtdistanz += distanz

        print(f"\nFahrzeug {fahrzeug + 1}:")
        print("  " + " -> ".join(stationen))
        print(f"  Schichtzeit {endzeit} min | Fahrstrecke {distanz} km | "
              f"Fracht {fracht}/{kapazitaet[fahrzeug]} "
              f"({fracht / kapazitaet[fahrzeug] * 100:.0f} % Auslastung)")

    print("\n" + "-" * 84)
    print(f"Summe Schichtzeiten:   {gesamtzeit} min")
    print(f"Summe Fahrstrecken:    {gesamtdistanz} km")
    print(f"Transportierte Fracht: {gesamtfracht} von {sum(daten['bedarfe'])} Einheiten")
    assert gesamtfracht == sum(daten["bedarfe"]), "Nicht alle Kunden wurden beliefert!"
    print("Alle Kunden wurden innerhalb ihrer Zeitfenster beliefert.")
    print("=" * 84)


if __name__ == "__main__":
    loese_cvrptw()
