import agx
import agxCollide
import agxUtil
import agxSDK
import agxRender
import agxVehicle
import agxOSG
from agxPythonModules.utils.environment import simulation, init_app, application, root
from math import cos, pi


# Help method to create a box car with GTires and wheel joints.
def createBoxCarWheelJoint(tire_material: agx.Material = None,
                           min_reference_speed=None,
                           clearance=0.4,
                           chassis_mass=1500,
                           chassis_cm_position=None,
                           use_fine_tire_slicing=False) -> agxSDK.Assembly:
    assembly = agxSDK.Assembly()

    # Parameters describing the car size. These are also
    # use when positioning the wheels
    car_width = 2.4
    car_length = 4.0
    car_height = 1.0

    car_half_lengths = agx.Vec3(car_width, car_length, car_height) * 0.5

    wheel_radius = 0.4
    wheel_width = 0.2

    # Create car body
    chassis_body = agx.RigidBody('chassis')
    chassis_geom = agxCollide.Geometry(agxCollide.Box(car_half_lengths))
    chassis_body.add(chassis_geom)
    chassis_body.getMassProperties().setMass(chassis_mass)
    if chassis_cm_position is not None:
        chassis_body.setCmPosition(chassis_cm_position)

    assembly.add(chassis_body)

    chassis_frame = chassis_body.getFrame()

    steering_axis = chassis_frame.transformVectorToWorld(agx.Vec3.Z_AXIS())
    wheel_axis = chassis_frame.transformVectorToWorld(agx.Vec3.X_AXIS())

    # Create visualization
    chassis_node = agxOSG.createVisual(chassis_body, root())
    agxOSG.setDiffuseColor(chassis_node, agxRender.Color.Red())
    agxOSG.forceWireFrameModeOn(chassis_node)

    tires = []
    wheel_joints = []

    # Y - front or back wheels
    # X - Which side of the car
    for y in [1, -1]:
        for x in [1, -1]:
            tire_shape = agxCollide.TireShape(wheel_radius, wheel_width)
            if use_fine_tire_slicing:
                slicing = agxCollide.TireSlicingParameters()
                slicing.numAxialSlices = 3
                slicing.numRadialSlicesPerIntersection = 5
                slicing.minChordLengthForRadialSlicing = 0.025
                tire_shape.setSlicingParameters(slicing)
            tire_geom = agxCollide.Geometry(tire_shape)
            tire_body = agx.RigidBody(tire_geom)
            tire_geom.setMaterial(tire_material)

            # Position for wheel
            wheel_position = agx.Vec3(x * (car_half_lengths[0] + wheel_width),
                                      y * (car_half_lengths[1] - 0.4),
                                      -car_half_lengths[2] - clearance + wheel_radius)

            tire_body.setPosition(chassis_frame.transformPointToWorld(wheel_position))
            tire_body.setRotation(agx.Quat(agx.Vec3.Y_AXIS(), agx.Vec3.X_AXIS()))

            # Adjust the tire mass, do not want the default generated value
            tire_body.getMassProperties().setMass(10)

            # Create GTire before the WheelJoint. GTire creates a non-spinning
            # carrier at the tire-body transform.
            tire = agxVehicle.GTire(tire_body)
            tire.setEnableRendering(True)
            if min_reference_speed is not None:
                tire.setGTireParameters(createGTireParameters(min_reference_speed))
            tire_carrier = tire.getTireCarrier()
            tires.append(tire)
            assembly.add(tire)

            # Create visuals
            tire_node = agxOSG.createVisual(tire_body, root())
            agxOSG.setDiffuseColor(tire_node, agxRender.Color.Black())

            # Attach the wheel and chassis with a WheelJoint:
            point_chassis_coords = agx.Vec3(x * (car_half_lengths[0]),
                                            y * (car_half_lengths[1] - 2 * wheel_width),
                                            -car_half_lengths[2] - clearance + wheel_radius)

            point_world = chassis_frame.transformPointToWorld(point_chassis_coords)

            # The WheelJoint provides only suspension and steering in this setup.
            # Wheel rotation and propulsion are handled by the GTire carcass hinge.
            wheel_joint_frame = agxVehicle.WheelJointFrame(point_world, wheel_axis, steering_axis)
            wheel_joint = agxVehicle.WheelJoint(wheel_joint_frame, tire_carrier, chassis_body)

            wheel_joints.append(wheel_joint)
            assembly.add(wheel_joint)

    # Ackerman steering contraint to front wheels
    front_wheels = wheel_joints[:2]
    steering_constraint = agxVehicle.Ackermann(front_wheels[1],
                                               front_wheels[0])
    simulation().add(steering_constraint)

    for i, wheel_joint in enumerate(wheel_joints):
        agxUtil.setEnableCollisions(chassis_body, tires[i].getTireBody(), False)
        wheel_joint.getLock1D(agxVehicle.WheelJoint.STEERING).setEnable(i >= 2)

        # WheelJoint is used only for suspension and steering. Tire rotation
        # and propulsion belong to the GTire carcass hinge.
        wheel_joint.getLock1D(agxVehicle.WheelJoint.WHEEL).setEnable(True)

        # Set the range for the suspension. That is the max travel distance
        wheel_joint.getRange1D(agxVehicle.WheelJoint.SUSPENSION).setEnable(True)
        wheel_joint.getRange1D(agxVehicle.WheelJoint.SUSPENSION).setRange(-0.1, 0.1)

        # Setup the suspension, stiffness and damping
        spring_constant = 1E5
        spring_damping = 4E3
        lock = wheel_joint.getLock1D(agxVehicle.WheelJoint.SUSPENSION)
        lock.setEnable(True)
        lock.setCompliance(agxUtil.convertSpringConstantToCompliance(spring_constant))
        lock.setDamping(agxUtil.convertDampingCoefficientToSpookDamping(
            spring_damping, spring_constant))

    return assembly, tires, wheel_joints, steering_constraint


def createGTireParameters(min_reference_speed):
    """Return the parameter set used by the traversal test scenes."""
    parameters = agxVehicle.GTireParameters()
    parameters.nominalInflationPressure = 220000
    parameters.inflationPressure = 220000
    parameters.verticalStiffness = 209651
    parameters.dKz_dP = 0.7098
    parameters.verticalDamping = 150
    parameters.longitudinalStiffness = 358066
    parameters.dKx_dP_lin = 0.17504
    parameters.dKx_dP_quad = 0
    parameters.lateralStiffness = 358066
    parameters.dKy_dP = 0.16365
    parameters.yawStiffness = 6000
    parameters.muStatic = 1.0
    parameters.muDynamic = 1.0
    parameters.rollingResistance = 0.00702
    parameters.rollingResistance_v = 0.001515
    parameters.rollingResistance_v4 = 8.514e-5
    parameters.rollingResistance_Fz = 0
    parameters.rollingResistance_P = 0
    parameters.minReferenceSpeed = min_reference_speed
    return parameters


class VehicleKeyboardListener(agxSDK.GuiEventListener):
    def __init__(self):
        super().__init__(agxSDK.GuiEventListener.KEYBOARD)
        self.up_pressed = False
        self.down_pressed = False
        self.left_pressed = False
        self.right_pressed = False

    def keyboard(self, key, modifier, x, y, keydown):
        if key == self.KEY_Up:
            self.up_pressed = bool(keydown)
        elif key == self.KEY_Down:
            self.down_pressed = bool(keydown)
        elif key == self.KEY_Left:
            self.left_pressed = bool(keydown)
        elif key == self.KEY_Right:
            self.right_pressed = bool(keydown)
        else:
            return False

        return True


class VehicleController(agxSDK.StepEventListener):
    def __init__(self, wheel_joints, carcass_hinges, steering_constraint, keyboard,
                 drive_force_limit=1E3):
        super().__init__()
        self.steering_constraint = steering_constraint
        self.keyboard = keyboard
        self.drive_motors = []
        self.drive_force_limit = drive_force_limit
        self.reverse_drive_force_scale = 0.5
        self.coasting_time = 3.0
        self.time_since_drive_release = self.coasting_time
        self.previous_drive_direction = 0

        for wheel_joint in wheel_joints:
            wheel_joint.getMotor1D(agxVehicle.WheelJoint.STEERING).setEnable(False)

        for i, carcass_hinge in enumerate(carcass_hinges):
            motor = carcass_hinge.getMotor1D()
            motor.setEnable(i < 2)
            if i < 2:
                motor.setForceRange(agx.RangeReal(drive_force_limit))
                self.drive_motors.append(motor)

    def pre(self, time):
        steer_direction = int(self.keyboard.left_pressed) - int(self.keyboard.right_pressed)
        if steer_direction != 0:
            steering_angle = self.steering_constraint.getSteeringAngle()
            self.steering_constraint.setSteeringAngle(
                steering_angle + steer_direction * 1.5 * simulation().getTimeStep())

        drive_direction = int(self.keyboard.up_pressed) - int(self.keyboard.down_pressed)
        if drive_direction != 0:
            self.time_since_drive_release = 0.0
            # Reverse uses less available torque to avoid excessive pitch and
            # rear-axle lift during backward acceleration.
            force_limit = (self.drive_force_limit if drive_direction > 0 else
                           self.reverse_drive_force_scale * self.drive_force_limit)
            for motor in self.drive_motors:
                motor.setSpeed(-drive_direction * 20)
                motor.setForceRange(agx.RangeReal(force_limit))
                motor.setEnable(True)
        else:
            if self.previous_drive_direction != 0:
                self.time_since_drive_release = 0.0

            if (self.previous_drive_direction != 0 or
                    self.time_since_drive_release < self.coasting_time):
                # Disable the motors for three seconds after key release so
                # that the vehicle can roll freely before braking.
                for motor in self.drive_motors:
                    motor.setEnable(False)
                self.time_since_drive_release += simulation().getTimeStep()
            else:
                # After three seconds of free rolling, re-enable the motors with zero
                # target speed. The enabled velocity motors then act as brakes.
                for motor in self.drive_motors:
                    motor.setSpeed(0)
                    motor.setForceRange(agx.RangeReal(self.drive_force_limit))
                    motor.setEnable(True)

        self.previous_drive_direction = drive_direction


def hfIndexToPosition(index, resolution, size):
    return (index / float(resolution - 1) - 0.5) * size


def createProfileHeightField(height_function,
                             terrain_length=80.0,
                             terrain_width=16.0,
                             element_size=0.4):
    """Create a height field by extruding a longitudinal height profile."""
    x_resolution = int(terrain_length / element_size) + 1
    y_resolution = int(terrain_width / element_size) + 1
    height_field = agxCollide.HeightField(
        x_resolution, y_resolution, terrain_length, terrain_width)

    for i in range(x_resolution):
        height = height_function(hfIndexToPosition(i, x_resolution, terrain_length))
        for j in range(y_resolution):
            height_field.setHeight(i, j, height)

    return height_field


def speedBumpHeight(x, bump_center=0.0, bump_height=0.25, bump_length=1.0):
    """Cosine-shaped bump with continuous height and slope at both ends."""
    half_length = 0.5 * bump_length
    if abs(x - bump_center) > half_length:
        return 0.0
    return 0.5 * bump_height * (
        1.0 + cos(2.0 * pi * (x - bump_center) / bump_length))


def slopeHeight(x,
                climb_start=-12.0,
                climb_end=0.0,
                descent_start=4.0,
                descent_end=16.0,
                slope_height=2.5):
    """Ramp profile with a flat top between its climb and descent."""
    if x < climb_start:
        return 0.0
    if x < climb_end:
        return slope_height * (x - climb_start) / (climb_end - climb_start)
    if x < descent_start:
        return slope_height
    if x < descent_end:
        return slope_height * (1.0 - (x - descent_start) / (descent_end - descent_start))
    return 0.0


def buildProfileTraversalTest(height_function, material_name,
                              start_x, min_reference_speed, title):
    height_field = createProfileHeightField(height_function)
    ground_geometry = agxCollide.Geometry(height_field)
    ground_material = agx.Material(material_name)
    ground_geometry.setMaterial(ground_material)
    ground = agx.RigidBody(ground_geometry)
    ground.setMotionControl(agx.RigidBody.STATIC)
    simulation().add(ground)
    agxOSG.createVisual(ground, root())

    tire_material = agx.Material('TireMaterial')

    car, tires, wheel_joints, steering = createBoxCarWheelJoint(
        tire_material,
        min_reference_speed=min_reference_speed,
        clearance=0.95,
        chassis_mass=3000,
        chassis_cm_position=agx.Vec3(0, 0, 0.5),
        use_fine_tire_slicing=True)
    car.setPosition(start_x, 0.0, 1.6 + height_function(start_x))
    car.setRotation(agx.EulerAngles(0, 0, -0.5 * agx.PI))
    simulation().add(car)

    keyboard = VehicleKeyboardListener()
    simulation().add(keyboard)
    simulation().add(VehicleController(
        wheel_joints,
        [tire.getCarcassHinge() for tire in tires],
        steering,
        keyboard,
        drive_force_limit=5E3))
    simulation().setPreIntegratePositions(True)

    application().setEnableDebugRenderer(True)
    decorator = application().getSceneDecorator()
    decorator.setText(0, title, agxRender.Color.Yellow())
    decorator.setText(1, 'Drive: UP/DOWN arrow keys')
    decorator.setText(2, 'Steer: LEFT/RIGHT arrow keys')


def speed_bump_traversal_test():
    """Drive across a speed bump to demonstrate vertical tire response."""
    buildProfileTraversalTest(
        speedBumpHeight, 'speed_bump_ground', -16.0, 1E-5,
        'GTire speed-bump traversal')


def slope_climb_descent_test():
    """Climb the ramp, cross its flat top, and descend the other side."""
    buildProfileTraversalTest(
        slopeHeight, 'slope_ground', -16.0, 1E-1,
        'GTire slope climb and descent')


# This tutorial will show how you can use the tire model in a more complex setting,
# a box car in this case.
def box_car_tutorial():
    # A height field as ground that the car can drive upon.
    ground_material = agx.Material('GroundMaterial')
    hf = agxUtil.HeightFieldGenerator.createHeightFieldFromFile("textures/terrain_test/terrain_height.png", 90, 90, -10, 20)
    assert hf
    ground_geo = agxCollide.Geometry(hf)
    ground = agx.RigidBody(ground_geo)
    ground.setMotionControl(agx.RigidBody.STATIC)
    ground_material = agx.Material('ground')
    ground_geo.setMaterial(ground_material)
    simulation().add(ground)
    agxOSG.createVisual(ground, root())

    tire_material = agx.Material('TireMaterial')

    # Create GTires first, then attach their carriers to the chassis with
    # suspension-and-steering-only WheelJoints.
    car, tires, wheel_joints, steering_constraint = createBoxCarWheelJoint(tire_material)

    # Move the car up a bit so it is above the height field.
    car.setPosition(0, 0, 3.6)
    car.setRotation(agx.EulerAngles(0, 0, agx.PI))
    simulation().add(car)

    carcass_hinges = [tire.getCarcassHinge() for tire in tires]
    keyboard = VehicleKeyboardListener()
    simulation().add(keyboard)
    simulation().add(VehicleController(
        wheel_joints, carcass_hinges, steering_constraint, keyboard))


# The basics about GTire and how to create it.
def basic_tutorial():
    # To create a completely new tire, just call the method with the wanted dimensions
    # and add it to the simulation
    tire_radius = 0.5
    tire_width = 0.3
    tire = agxVehicle.GTire(tire_radius, tire_width)
    simulation().add(tire)
    tire.setPosition(0, 0, 0.5)

    # Next, let's create a height field for the tire to collide with
    hf = agxCollide.HeightField(10, 10, 5, 5, 10)
    ground_geo = agxCollide.Geometry(hf)
    ground = agx.RigidBody(ground_geo)
    ground.setMotionControl(agx.RigidBody.STATIC)
    simulation().add(ground)

    # The materials identify the tire-ground contact pair. GTire obtains its
    # operating coefficients from GTireParameters.
    ground_material = agx.Material('ground')
    ground_geo.setMaterial(ground_material)
    tire_material = agx.Material('tire')
    tire.setMaterial(tire_material)

    # Enable debug rendering
    application().setEnableDebugRenderer(True)


def addRemainingScenes(app):
    '''
    Add all scenes so we can switch between them using 1..n keys or +/- keys.
    '''
    scriptFileName = app.getArguments().getArgumentName(1)
    scriptFileName = scriptFileName.replace('agxscene:', '')

    def addScene(name):
        sceneKey = app.getNumScenes() + 1
        app.addScene(scriptFileName, name, ord(ascii(sceneKey)), True)

    addScene("basic_tutorial")
    addScene("speed_bump_traversal_test")
    addScene("slope_climb_descent_test")


def buildScene():
    """
    Entry point when running this script using agxViewer.
    """

    if application().getNumScenes() == 1:
        addRemainingScenes(application())

    box_car_tutorial()


# Entry point when this script is started with python executable
init = init_app(name=__name__,
                scenes=[
                    (buildScene, '1'),
                    (basic_tutorial, '2'),
                    (speed_bump_traversal_test, '3'),
                    (slope_climb_descent_test, '4'),
                ],
                autoStepping=True,  # Default: False
                )
