import numpy as np
from scipy.optimize import minimize, fsolve
import matplotlib.pyplot as plt
from mpl_toolkits.mplot3d import Axes3D
import sys

def equation(delta_o_guess,d,w,beta,delta_i,tag):
    lhs = (w-2*d*np.sin(beta))**2
    
    if tag == "rear":
        #print("Rear trapeize")
        rhs = d**2*(np.cos(delta_o_guess-beta)-np.cos(delta_i+beta))**2+(w+d*(np.sin(delta_o_guess-beta)-np.sin(delta_i+beta)))**2
    elif tag == "front":
        #print("front trapeize")
        rhs = d**2*(np.cos(beta-delta_i)-np.cos(beta+delta_o_guess))**2 + (w-d*(np.sin(beta-delta_i)+np.sin(beta+delta_o_guess)))**2
    else:
        print(f"[***ERROR] geometry type {tag} unknown. Script aborted")
        sys.exit()

    #print("lhs, rhs",lhs,rhs)
    
    return lhs - rhs


def solve_delta_o_trap(delta_o_ak,d, w, beta, delta_i,tag):
    """
    Solve for delta_o_trap from the trapezoidal equation.
        """

    
    # Try to solve the equation
    try:
        # Initial guess for delta_o_trap (could be delta_o_akerman as a starting point)
        initial_guess = delta_o_ak
        #print("initial guess for delta_o: ",initial_guess)
        solution = fsolve(equation, initial_guess, args=(d,w,beta,delta_i,tag), full_output=True)
        #print(solution)
        delta_o_trap = solution[0][:]
        info = solution[1]
        
        # Check if solution converged
        if info['fvec'][0]**2 < 1e-10:
            #print(solution[3])
            return delta_o_trap
        else:
            return np.nan
    except:
        return np.nan

def calculate_error(params, w, l, delta_i_array,tag):
    """
    Calculate the RMS error for given d and beta values across all delta_i values.
    
    Parameters:
    - params: [d, beta] array
    - w, l: problem parameters
    - delta_i_array: array of delta_i values to evaluate
    """
    d, beta = params
    
    errors = []
    
    for delta_i in delta_i_array:
        # Calculate delta_o_ak
        delta_o_ak = find_delta_o_akerman(w,l,delta_i)
        
        # Solve for delta_o_trap
        delta_o_trap = solve_delta_o_trap(delta_o_ak,d, w, beta, delta_i,tag=tag)
        
        if np.isnan(delta_o_trap):
            # If no valid solution, add a large penalty
            errors.append(1e10)
        else:
            # Calculate absolute error
            error = np.abs(delta_o_ak - delta_o_trap)
            errors.append(error)
    
    # Calculate RMS error
    rms_error = np.sqrt(np.mean(np.array(errors)**2))
    
    return rms_error
def find_delta_o_akerman(w,l,delta_i):
    return np.arctan(1/( w / l + 1 / np.tan(delta_i)))

def find_optimal_d_and_beta(w, l, delta_i_max, num_points=50, 
                            d_initial=None, beta_initial=None,
                            d_bounds=None, beta_bounds=None,tag="rear"):
    """
    Find the optimal values of d and beta that minimize the RMS error.
    
    Parameters:
    - w, l: problem parameters
    - delta_i_max: maximum angle (in radians)
    - num_points: number of points to sample in delta_i range
    - d_initial: initial guess for d (if None, will be estimated)
    - beta_initial: initial guess for beta (if None, will be estimated)
    - d_bounds: tuple (d_min, d_max) for d bounds
    - beta_bounds: tuple (beta_min, beta_max) for beta bounds
    """
    # Create array of delta_i values in radians (avoid 0 due to cotangent)
    delta_i_array = np.linspace(np.radians(1.0), delta_i_max, num_points)
    
    # Set initial guesses if not provided
    if d_initial is None:
        d_initial = w * 2
    if beta_initial is None:
        beta_initial = np.radians(30)
    
    # Set bounds if not provided
    if d_bounds is None:
        d_bounds = (w / 10, w * 20)
    if beta_bounds is None:
        beta_bounds = (np.radians(1), np.radians(89))
    
    # Initial parameters
    initial_params = [d_initial, beta_initial]
    
    # Bounds for optimization
    bounds = [d_bounds, beta_bounds]
    
    # Define objective function
    def objective(params):
        return calculate_error(params, w, l, delta_i_array,tag=tag)
    
    # Optimize using multiple methods for robustness
    print("Starting optimization ...")
    print("Initial parameters [d, beta]: ",initial_params)
    print("Optimization bounds [(d_min,d_max),(beta_min,beta_max)]: ", bounds)
    
    # Try with L-BFGS-B method
    result = minimize(objective, initial_params, method='L-BFGS-B', 
                     bounds=bounds, options={'ftol': 1e-12, 'maxiter': 1000})
    
    optimal_d, optimal_beta = result.x
    min_rms_error = result.fun
    
    print(f"Optimization converged: {result.success}")
    print(f"Number of iterations: {result.nit}")
    
    return optimal_d, optimal_beta, min_rms_error, delta_i_array

# Example usage
if __name__ == "__main__":
    # Define parameters (angles in radians)
    w = 1.1 # Track in [m]
    l = 1.6 # wheelbase in [m]
    delta_i_max = np.radians(30)  # maximal steering angle of inner wheel
    beta_initial = np.radians(10) # initial guess value for trapeize angle
    beta_min = np.radians(1.0) # minimal trapeize angle
    beta_max = np.radians(60.0) # maximal trapeize angle
    d_initial = 0.2 # initial guess value for arm length in [m]
    d_min = 0.1 # minimal arm length in [m]
    d_max = 0.5 # maximal arm length in [m]
    num_points = 50
    tag = "rear" # "rear" for a mecanism located rear of the wheels. "front" for a mechanism located in front of the wheels.
    plot_option = "3D" # "contour" for contour plot, "3D" for a 3D plot, "minRMS" to plot the min RMS error versus de arm length d
    output_file_name = "rear-optimum.txt"

    # delta_i_array = np.linspace(np.radians(1.0), delta_i_max, num_points)
    # #print("delta_i_array: ", np.degrees(delta_i_array))
    # delta_o_ak = find_delta_o_akerman(w,l,delta_i_array)
    # #print("delta_o_ak: ", np.degrees(delta_o_ak))
    # delta_o_trap = solve_delta_o_trap(delta_o_ak,d_initial,w,beta_initial,delta_i_array)
    # print("delta_i: ",delta_i_array)
    # print("delta_o_ak: ",delta_o_ak)
    # print("delta_o_trap: ", delta_o_trap)
    
    # Find optimal d and beta
    optimal_d, optimal_beta, min_rms_error, delta_i_array = find_optimal_d_and_beta(
        w, l, delta_i_max, num_points=num_points,d_initial=d_initial, beta_initial=beta_initial,d_bounds=(d_min,d_max), beta_bounds=(beta_min,beta_max),tag=tag)
    
    # Prepare the output text
    output_text = f"\n{'='*50}\n"
    output_text += f"OPTIMIZATION RESULTS\n"
    output_text += f"{'='*50}\n"
    output_text += f"Optimal d: {optimal_d:.6f}\n"
    output_text += f"Optimal beta: {np.degrees(optimal_beta):.6f} degrees ({optimal_beta:.6f} radians)\n"
    output_text += f"Minimum angle RMS error between Akerman and trapeizoidal mechanism: {np.degrees(min_rms_error):.6e} [deg]\n"
    output_text += f"{'='*50}\n\n"

    # Print to console
    print(output_text)

    # Save to file
    with open(output_file_name, 'w') as f:
        f.write(output_text)

    print(f"Results saved to {output_file_name}")
    
    # Visualization    
    delta_o_ak_array = find_delta_o_akerman(w,l,delta_i_array)
    delta_o_trap_array = solve_delta_o_trap(delta_o_ak_array,optimal_d, w, optimal_beta, delta_i_array,tag=tag)    

    # Plot comparison of delta_o functions
    plt.figure(figsize=(12, 5))
    
    plt.subplot(1, 2, 1)
    plt.plot(np.degrees(delta_i_array), np.degrees(delta_o_ak_array), 'b-', label='Akerman', linewidth=2)
    plt.plot(np.degrees(delta_i_array), np.degrees(delta_o_trap_array), 'r--', label='Trapezoidal mechanism', linewidth=2)
    plt.xlabel(r'$\delta_i$ [deg]', fontsize=12)
    plt.ylabel(r'$\delta_o$ [deg]', fontsize=12)
    plt.title(f'Comparison of optimized trapezoidal mechanism with Akerman angle\nd={optimal_d:.4f} m, β={np.degrees(optimal_beta):.2f} [deg]', fontsize=12)
    plt.legend(fontsize=10)
    plt.grid(True, alpha=0.3)
    
    # Plot error
    plt.subplot(1, 2, 2)
    errors = (np.array(delta_o_trap_array)-np.array(delta_o_ak_array))
    plt.plot(np.degrees(delta_i_array), np.degrees(errors), 'g-', linewidth=2)
    plt.xlabel(r'$\delta_i$ [deg]', fontsize=12)
    plt.ylabel(r'$\delta_o-\delta_{o,Akerman}$ [deg]', fontsize=12)
    plt.title(f'Error between Trapezoidal and Akerman\nRMS = {np.degrees(min_rms_error):.6e} [deg]', fontsize=12)
    plt.grid(True, alpha=0.3)
    
    plt.tight_layout()
    plt.show()
    
    # Optional: Create a contour plot showing the error landscape
    print("Generating error 3D plots...")
    
    # Create a grid of d and beta values
    d_range = np.linspace(optimal_d * 0.9, optimal_d /0.9, 30)
    beta_range = np.linspace(optimal_beta * 0.9, optimal_beta /0.9, 30)
    D, B = np.meshgrid(d_range, beta_range)
    
  
    


    # # Calculate error for each combination
    Z = np.zeros_like(D)
    for i in range(D.shape[0]):
        for j in range(D.shape[1]):
            Z[i, j] = calculate_error([D[i, j], B[i, j]], w, l, delta_i_array,tag=tag)

    if plot_option == "3D":
        fig = plt.figure(figsize=(12, 9))
        ax = fig.add_subplot(111, projection='3d')

        # Create 3D surface plot
        surf = ax.plot_surface(D, np.degrees(B), np.degrees(Z), cmap='viridis', alpha=0.8, 
                            edgecolor='none', antialiased=True)

        # Add colorbar
        cbar = fig.colorbar(surf, ax=ax, shrink=0.5, aspect=5)
        cbar.set_label(r'$\varepsilon_{RMS}(\delta_{o.Akerman}-\delta_{o.trapeize})$ [deg]', fontsize=12)

        # Mark the optimal point
        ax.scatter(optimal_d, np.degrees(optimal_beta), np.degrees(min_rms_error), 
                color='red', s=200, marker='*', label='Optimal point', 
                edgecolors='black', linewidths=2)

        # Labels and title
        ax.set_xlabel('d [m]', fontsize=12, labelpad=10)
        ax.set_ylabel(r'$\beta$ [deg]', fontsize=12, labelpad=10)
        ax.set_zlabel(r'$\varepsilon_{RMS}(\delta_{o.Akerman}-\delta_{o.trapeize})$ [deg]', fontsize=12, labelpad=10)
        ax.set_title(r'$\delta_o$ Error Landscape Akerman vs Trapeizoidal: RMS Error vs d and β', fontsize=14, pad=20)

        # Add legend
        ax.legend(fontsize=10, loc='upper left')

        # Adjust viewing angle for better visualization
        ax.view_init(elev=25, azim=45)

        plt.tight_layout()
        plt.show()
    elif plot_option == "contour":
        plt.figure(figsize=(10, 8))
        contour = plt.contour(D, np.degrees(B), Z, levels=20, cmap='viridis')
        plt.colorbar(contour, label='RMS Error')
        plt.plot(optimal_d, np.degrees(optimal_beta), 'r*', markersize=20, label='Optimal point')
        plt.xlabel('d [m]', fontsize=12)
        plt.ylabel(r'$\beta$ [deg]', fontsize=12)
        plt.title(r'$\delta_o$ Error Landscape Akerman vs Trapeizoidal: RMS Error vs d and β', fontsize=14)
        plt.legend(fontsize=10)
        plt.grid(True, alpha=0.3)
        plt.tight_layout()
        plt.show()
    elif plot_option == "minRMS":
        # Find minimum Z value along the B axis and the corresponding indices
        Z_min_along_B = np.min(Z, axis=0)
        beta_indices_at_min = np.argmin(Z, axis=0)

        # Get the d and beta values
        d_values = D[0, :]
        beta_values_at_min = B[beta_indices_at_min, np.arange(len(beta_indices_at_min))]

        # Create subplot with two plots
        fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(10, 10))

        # Plot 1: min(Z) vs d
        ax1.plot(d_values, Z_min_along_B, 'b-', linewidth=2, label='min(Z) over β')
        ax1.set_xlabel('d [m]', fontsize=12)
        ax1.set_ylabel(r'$\min_{\beta}(\varepsilon_{\mathrm{RMS}})$', fontsize=12)
        ax1.set_title(r'Minimum RMS Error vs d (minimized over $\beta$)', fontsize=14)
        ax1.grid(True, alpha=0.3)
        ax1.legend(fontsize=10)

        # Plot 2: Optimal beta vs d
        ax2.plot(d_values, np.degrees(beta_values_at_min), 'r-', linewidth=2, label='Optimal β')
        ax2.set_xlabel('d [m]', fontsize=12)
        ax2.set_ylabel(r'Optimal $\beta$ [degrees]', fontsize=12)
        ax2.set_title(r'Optimal $\beta$ vs d (at minimum error)', fontsize=14)
        ax2.grid(True, alpha=0.3)
        ax2.legend(fontsize=10)

        plt.tight_layout()
        plt.show()

        # Print overall optimal values
        optimal_idx = np.argmin(Z_min_along_B)
        optimal_d_from_slice = d_values[optimal_idx]
        optimal_beta_from_slice = beta_values_at_min[optimal_idx]
        optimal_error_from_slice = Z_min_along_B[optimal_idx]

        print(f"\n{'='*50}")
        print(f"Optimal d: {optimal_d_from_slice:.6f} m")
        print(f"Optimal β: {np.degrees(optimal_beta_from_slice):.2f} degrees")
        print(f"Minimum error: {optimal_error_from_slice:.6e}")
        print(f"{'='*50}\n")
    else:
        print("[***ERROR] plot option unkown.")