! Copyright (c) 2022-2026 Jason Christopherson
! SPDX-License-Identifier: MIT
!
! Permission is hereby granted, free of charge, to any person obtaining a copy
! of this software and associated documentation files (the "Software"), to deal
! in the Software without restriction, including without limitation the rights
! to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
! copies of the Software, and to permit persons to whom the Software is
! furnished to do so, subject to the following conditions:
!
! The Software is provided "as is", without warranty of any kind, express or
! implied, including but not limited to the warranties of merchantability,
! fitness for a particular purpose and noninfringement.
module dynamics_variational_integrators
    !! Provides a maximal-coordinate variational integrator for constrained
    !! rigid-body systems. The implementation follows the discrete
    !! Euler-Lagrange equations and graph factorization described by Brudigam,
    !! Sosnowski, Manchester, and Hirche (2023).
    !!
    !! Each body contributes three world-frame translational velocities and
    !! three body-frame angular velocities to the Newton system. Holonomic
    !! equality constraints contribute scalar Lagrange multipliers. The dense
    !! solver is useful as a reference implementation; the graph-factorized
    !! solver eliminates six-variable body nodes and scalar constraint nodes
    !! according to the sparsity graph of the Newton matrix.
    use iso_fortran_env, only : int32, real64
    use linalg, only : lu_factor, solve_lu
    use dynamics_error_handling, only : DYN_ARRAY_SIZE_ERROR, &
        DYN_CONVERGENCE_ERROR, DYN_INVALID_INPUT_ERROR
    use dynamics_quaternions, only : quaternion, operator(*), abs, inverse
    use dynamics_rigid_bodies, only : rigid_body
    use dynamics_helper, only : cross_product
    implicit none
    private

    public :: variational_state
    public :: variational_integrator_info
    public :: variational_integrator
    public :: variational_integrator_settings
    public :: variational_force
    public :: variational_constraint
    public :: variational_constraint_jacobian
    public :: initialize_variational_state
    public :: VI_DENSE_SOLVER
    public :: VI_GRAPH_FACTORIZED_SOLVER
    public :: VI_FORCE_LEFT_ENDPOINT
    public :: VI_FORCE_IMPLICIT_ENDPOINT
    public :: VI_FORCE_MIDPOINT

    integer(int32), parameter :: VI_DENSE_SOLVER = 1
        !! Selects conventional dense LU factorization of the Newton matrix.
    integer(int32), parameter :: VI_GRAPH_FACTORIZED_SOLVER = 2
        !! Selects graph-ordered block LDU factorization of the Newton matrix,
        !! including graph fill generated by Schur-complement updates.
    integer(int32), parameter :: VI_FORCE_LEFT_ENDPOINT = 1
        !! Evaluates applied loads at the current state, explicitly.
    integer(int32), parameter :: VI_FORCE_IMPLICIT_ENDPOINT = 2
        !! Evaluates applied loads at the trial next state inside Newton.
    integer(int32), parameter :: VI_FORCE_MIDPOINT = 3
        !! Evaluates applied loads at the interpolated midpoint state inside
        !! Newton. This selects force quadrature, not a different conservative
        !! state update or formal integrator order.

    type variational_state
        !! Defines the maximal-coordinate state of a collection of rigid
        !! bodies. Angular velocities and inertia tensors are expressed in each
        !! body's frame. Positions, translational velocities, forces, and
        !! orientations use the world frame.
        real(real64), allocatable, dimension(:,:) :: position
            !! Body center-of-mass positions, dimensioned 3-by-nbody.
        type(quaternion), allocatable, dimension(:) :: orientation
            !! Unit body-to-world orientation quaternions.
        real(real64), allocatable, dimension(:,:) :: velocity
            !! Center-of-mass velocities, dimensioned 3-by-nbody.
        real(real64), allocatable, dimension(:,:) :: angular_velocity
            !! Body-frame angular velocities, dimensioned 3-by-nbody.
        real(real64) :: time = 0.0d0
            !! The time associated with the state.
    end type

    type variational_integrator_info
        !! Reports convergence diagnostics for a step or trajectory solve.
        logical :: converged = .false.
            !! True when the requested solve completed successfully.
        integer(int32) :: iterations = 0
            !! Newton iterations used; aggregated across steps by solve.
        logical :: jacobian_singular = .false.
            !! True if any attempted Newton Jacobian was detected as singular,
            !! including one later regularized successfully.
    end type

    abstract interface
        subroutine variational_force(t, state, force, torque, args)
            !! Computes applied forces and body-frame torques at the selected
            !! force-evaluation state. Depending on the integrator settings,
            !! this is the current state, trial next state, or interpolated
            !! midpoint state. Velocity-dependent forces are reevaluated as
            !! Newton's trial state changes for the latter two modes.
            import :: real64, variational_state
            real(real64), intent(in) :: t
                !! The current simulation time.
            type(variational_state), intent(in) :: state
                !! The current maximal-coordinate state.
            real(real64), intent(out), dimension(:,:) :: force
                !! The 3-by-nbody world-frame force array.
            real(real64), intent(out), dimension(:,:) :: torque
                !! The 3-by-nbody body-frame torque array.
            class(*), intent(inout), optional :: args
                !! Optional user-supplied data.
        end subroutine

        subroutine variational_constraint(state, value, args)
            !! Evaluates the holonomic equality constraints for a supplied
            !! maximal-coordinate state.
            import :: real64, variational_state
            type(variational_state), intent(in) :: state
                !! The state at which to evaluate the constraints.
            real(real64), intent(out), dimension(:) :: value
                !! The constraint residual vector, which is zero for a
                !! constraint-compatible state.
            class(*), intent(inout), optional :: args
                !! Optional user-supplied data.
        end subroutine

        subroutine variational_constraint_jacobian(state, jacobian, args)
            !! Computes the reduced maximal-coordinate constraint Jacobian.
            !! Translational columns are ordinary position derivatives.
            !! Rotational columns use the local quaternion variation
            !! 
            !! $$q(\epsilon)=q\left(\sqrt{1-\epsilon^T\epsilon},\epsilon\right).$$
            import :: real64, variational_state
            type(variational_state), intent(in) :: state
                !! The state at which to evaluate the Jacobian.
            real(real64), intent(out), dimension(:,:) :: jacobian
                !! The nconstraint-by-(6*nbody) Jacobian. Columns are ordered
                !! [dx, quaternion-vector variation] for each body.
            class(*), intent(inout), optional :: args
                !! Optional user-supplied data.
        end subroutine
    end interface

    type variational_integrator_settings
        !! Defines numerical settings for the nonlinear variational step.
        real(real64) :: tolerance = 1.0d-10
            !! The Euclidean residual tolerance used to terminate Newton's
            !! method.
        real(real64) :: finite_difference_step = 1.0d-7
            !! The relative forward-difference step used for numerical
            !! Jacobians.
        real(real64) :: constraint_translation_scale = 1.0d0
            !! The characteristic translation magnitude used as the absolute
            !! floor when perturbing position components in a numerical
            !! constraint Jacobian. This value has the same length units as
            !! the model positions.
        real(real64) :: constraint_rotation_scale = 1.0d0
            !! The characteristic dimensionless quaternion-tangent magnitude
            !! used when perturbing rotational components in a numerical
            !! constraint Jacobian.
        integer(int32) :: maximum_iterations = 50
            !! The maximum number of Newton iterations allowed per time step.
        integer(int32) :: maximum_line_search_iterations = 12
            !! The maximum number of step halvings allowed by the residual
            !! line search.
        integer(int32) :: linear_solver = VI_DENSE_SOLVER
            !! The linear solver used for each Newton correction. Valid values
            !! are VI_DENSE_SOLVER and VI_GRAPH_FACTORIZED_SOLVER.
        integer(int32) :: force_evaluation = VI_FORCE_LEFT_ENDPOINT
            !! The applied-load evaluation point, used for every force and
            !! torque callback, including springs, dampers, and user loads.
            !! VI_FORCE_LEFT_ENDPOINT uses the current state explicitly;
            !! VI_FORCE_IMPLICIT_ENDPOINT uses the trial next state inside
            !! Newton; VI_FORCE_MIDPOINT uses the interpolated midpoint state
            !! inside Newton. The latter two modes reevaluate loads as the
            !! Newton trial changes. These modes change load quadrature, not
            !! the conservative state update or its formal order.
    end type

    type variational_integrator
        !! Implements the maximal-coordinate variational integrator in
        !! Brudigam et al. (2023), equations (18)-(20). The nonlinear system is
        !! solved with damped Newton iterations and either dense LU or the
        !! paper's graph-structured block factorization.
        type(variational_integrator_settings) :: settings
            !! Numerical settings used by the integrator.
    contains
        procedure, public :: step => vi_step
            !! Advances a maximal-coordinate rigid-body state by one step.
        procedure, public :: solve => vi_solve
            !! Computes the solution.
    end type

contains
! ------------------------------------------------------------------------------
subroutine initialize_variational_state(state, nbody)
    !! Allocates and initializes a maximal-coordinate state. Positions and
    !! velocities are set to zero, and every orientation is set to the identity
    !! quaternion.
    type(variational_state), intent(out) :: state
        !! The state to initialize.
    integer(int32), intent(in) :: nbody
        !! The number of rigid bodies represented by the state.

    ! Local Variables
    integer(int32) :: i

    ! Input Checking
    if (nbody < 1) error stop DYN_INVALID_INPUT_ERROR

    ! Initialization
    allocate(state%position(3, nbody), state%orientation(nbody), &
        state%velocity(3, nbody), state%angular_velocity(3, nbody))
    state%position = 0.0d0
    state%velocity = 0.0d0
    state%angular_velocity = 0.0d0
    state%time = 0.0d0
    do i = 1, nbody
        state%orientation(i) = quaternion([1.0d0, 0.0d0, 0.0d0, 0.0d0])
    end do
end subroutine

! ------------------------------------------------------------------------------
subroutine vi_step(this, bodies, state, dt, constraint_count, constraint, &
    force_function, constraint_jacobian, multipliers, args, info)
    !! Advances a rigid-body system by one fixed time step. The unknown vector
    !! contains the next translational and body-frame angular velocities,
    !! followed by the equality-constraint multipliers. Orientations are
    !! advanced with the unit quaternion increment
    !!
    !! $$q_{k+1}=q_k\left(\sqrt{1-\|h\omega_{k+1}/2\|^2},
    !! h\omega_{k+1}/2\right).$$
    class(variational_integrator), intent(in) :: this
        !! The variational integrator.
    type(rigid_body), intent(in), dimension(:) :: bodies
        !! The body mass and body-frame inertia properties.
    type(variational_state), intent(inout) :: state
        !! On input, the current state; on output, the converged next state.
        !! On recoverable failure, the state is unchanged.
    real(real64), intent(in) :: dt
        !! The positive fixed time step.
    integer(int32), intent(in), optional :: constraint_count
        !! The number of scalar equality constraints. The default is zero.
    procedure(variational_constraint), pointer, intent(in), optional :: constraint
        !! The holonomic equality-constraint callback. It is required when
        !! constraint_count is greater than zero.
    procedure(variational_force), pointer, intent(in), optional :: force_function
        !! The external force and torque callback. When omitted, all applied
        !! forces and torques are zero.
    procedure(variational_constraint_jacobian), pointer, intent(in), optional :: constraint_jacobian
        !! An optional analytic reduced constraint Jacobian. When omitted, the
        !! Jacobian is evaluated by finite differences.
    real(real64), allocatable, intent(out), optional, dimension(:) :: multipliers
        !! The converged Lagrange multipliers; unallocated on recoverable
        !! failure.
    class(*), intent(inout), optional :: args
        !! Optional user-supplied data forwarded to all callbacks.
    type(variational_integrator_info), intent(out), optional :: info
        !! Optional convergence diagnostics. When present, convergence failures
        !! are returned instead of terminating execution; when absent, the
        !! legacy error-stop behavior is retained.

    ! Local Variables
    integer(int32) :: i, iteration, iterations_used, line_iteration, &
        nbody, nconstraint, nvar
    integer(int32), allocatable, dimension(:) :: pivot
    real(real64) :: alpha, trial_norm, residual_norm
    logical :: jacobian_singular, graph_singular
    real(real64), allocatable, dimension(:) :: unknown, trial, residual, &
        trial_residual, delta, perturbed, perturbed_value, &
        base_constraint_value, constraint_value
    real(real64), allocatable, dimension(:,:) :: jacobian, lu, &
        constraint_gradient, applied_force, applied_torque
    type(variational_state) :: accepted_state, next_state, perturbed_state

    ! Input Checking
    nbody = size(bodies)
    nconstraint = 0
    if (present(constraint_count)) nconstraint = constraint_count
    call check_inputs(this%settings, bodies, state, dt, nconstraint, &
        present(constraint))
    jacobian_singular = .false.
    iterations_used = 0
    if (present(info)) then
        info%converged = .false.
        info%iterations = 0
        info%jacobian_singular = .false.
    end if

    ! Initialize the Newton unknown with the current velocities and zero
    ! equality-constraint impulses.
    nvar = 6 * nbody + nconstraint
    allocate(unknown(nvar), trial(nvar), residual(nvar), &
        trial_residual(nvar), jacobian(nvar, nvar), &
        perturbed(nvar), perturbed_value(nvar), &
        applied_force(3,nbody), applied_torque(3,nbody), &
        constraint_value(max(1,nconstraint)))
    next_state = state
    perturbed_state = state
    do i = 1, nbody
        unknown(6*i-5:6*i-3) = state%velocity(:,i)
        unknown(6*i-2:6*i) = state%angular_velocity(:,i)
    end do
    if (nconstraint > 0) unknown(6*nbody+1:nvar) = 0.0d0

    ! The reduced constraint Jacobian is evaluated at the current state and is
    ! constant throughout the Newton iterations for this discrete step.
    if (nconstraint > 0) then
        allocate(constraint_gradient(nconstraint, 6*nbody), &
            base_constraint_value(nconstraint))
        if (present(constraint_jacobian)) then
            call constraint_jacobian(state, constraint_gradient, args)
        else
            call finite_difference_constraint_gradient(constraint_gradient)
        end if
    end if

    ! Solve the discrete Euler-Lagrange equations and position-level
    ! constraints with damped Newton iterations.
    call evaluate_residual(unknown, residual)
    residual_norm = norm2(residual)
    do iteration = 1, this%settings%maximum_iterations
        if (residual_norm <= this%settings%tolerance) exit
        iterations_used = iteration
        call finite_difference_jacobian(unknown, residual, jacobian)
        select case (this%settings%linear_solver)
        case (VI_DENSE_SOLVER)
            call lu_factor(jacobian, ipvt = pivot, lu = lu)
            if (any(pivot == 0)) then
                jacobian_singular = .true.
                do i = 1, nvar
                    jacobian(i,i) = jacobian(i,i) + this%settings%tolerance
                end do
                call lu_factor(jacobian, ipvt = pivot, lu = lu)
                if (any(pivot == 0)) then
                    call convergence_failure(iterations_used, .true.)
                    return
                end if
            end if
            delta = solve_lu(lu, pivot, -residual)
        case (VI_GRAPH_FACTORIZED_SOLVER)
            delta = graph_factorized_solve(jacobian, -residual, nbody, &
                nconstraint, graph_singular)
            if (graph_singular) then
                call convergence_failure(iterations_used, .true.)
                return
            end if
        end select

        alpha = 1.0d0
        trial_norm = huge(1.0d0)
        do line_iteration = 0, this%settings%maximum_line_search_iterations
            trial = unknown + alpha * delta
            call evaluate_residual(trial, trial_residual)
            trial_norm = norm2(trial_residual)
            if (trial_norm < residual_norm) exit
            alpha = 0.5d0 * alpha
        end do
        if (trial_norm >= residual_norm) then
            call convergence_failure(iterations_used, jacobian_singular)
            return
        end if
        unknown = trial
        residual = trial_residual
        residual_norm = trial_norm
    end do
    if (residual_norm > this%settings%tolerance) then
        call convergence_failure(iterations_used, jacobian_singular)
        return
    end if

    ! Commit the converged state through a temporary to avoid aliasing the host
    ! state used by state_from_unknown.
    call state_from_unknown(unknown, accepted_state)
    accepted_state%time = state%time + dt
    state = accepted_state
    if (present(multipliers)) then
        allocate(multipliers(nconstraint))
        if (nconstraint > 0) multipliers = unknown(6*nbody+1:nvar)
    end if
    if (present(info)) then
        info%converged = .true.
        info%iterations = iterations_used
        info%jacobian_singular = jacobian_singular
    end if

contains
    subroutine convergence_failure(iterations, singular)
        integer(int32), intent(in) :: iterations
        logical, intent(in) :: singular

        if (present(info)) then
            info%converged = .false.
            info%iterations = iterations
            info%jacobian_singular = singular
        else
            error stop DYN_CONVERGENCE_ERROR
        end if
    end subroutine

    subroutine evaluate_residual(x, value)
        !! Evaluates the coupled discrete Euler-Lagrange and holonomic
        !! constraint residual for a Newton iterate.
        real(real64), intent(in), dimension(:) :: x
            !! The Newton unknown vector.
        real(real64), intent(out), dimension(:) :: value
            !! The residual vector.

        ! Local Variables
        type(variational_state) :: next_state, force_state
        type(quaternion) :: relative_increment, midpoint_increment
        integer(int32) :: body_index, first
        real(real64), dimension(3) :: omega1, omega2, momentum1, momentum2
        real(real64), dimension(3,3) :: rotation_start, rotation_end, &
            rotation_midpoint
        real(real64), dimension(4) :: increment_components
        real(real64) :: midpoint_scalar

        ! Construct the configuration implied by the trial velocities.
        call state_from_unknown(x, next_state)

        ! Evaluate applied loads at the configured quadrature point. Endpoint
        ! and midpoint modes make the force part of the Newton residual.
        applied_force = 0.0d0
        applied_torque = 0.0d0
        if (present(force_function)) then
            select case (this%settings%force_evaluation)
            case (VI_FORCE_LEFT_ENDPOINT)
                call force_function(state%time, state, applied_force, &
                    applied_torque, args)
            case (VI_FORCE_IMPLICIT_ENDPOINT)
                call force_function(next_state%time, next_state, applied_force, &
                    applied_torque, args)
            case (VI_FORCE_MIDPOINT)
                force_state = state
                force_state%time = state%time + 0.5d0 * dt
                force_state%position = 0.5d0 * &
                    (state%position + next_state%position)
                force_state%velocity = 0.5d0 * &
                    (state%velocity + next_state%velocity)
                do body_index = 1, nbody
                    relative_increment = inverse(state%orientation(body_index)) * &
                        next_state%orientation(body_index)
                    increment_components = relative_increment%to_array()
                    midpoint_scalar = sqrt(0.5d0 * &
                        max(0.0d0, min(2.0d0, &
                        1.0d0 + increment_components(1))))
                    if (midpoint_scalar > sqrt(tiny(1.0d0))) then
                        midpoint_increment = quaternion([midpoint_scalar, &
                            increment_components(2:4) / &
                            (2.0d0 * midpoint_scalar)])
                    else if (norm2(increment_components(2:4)) > &
                        sqrt(tiny(1.0d0))) then
                        midpoint_increment = quaternion([sqrt(0.5d0), &
                            increment_components(2:4) / &
                            norm2(increment_components(2:4)) * sqrt(0.5d0)])
                    else
                        midpoint_increment = &
                            quaternion([1.0d0, 0.0d0, 0.0d0, 0.0d0])
                    end if
                    force_state%orientation(body_index) = &
                        state%orientation(body_index) * midpoint_increment
                    call force_state%orientation(body_index)%normalize()
                    rotation_start = state%orientation(body_index)%to_matrix()
                    rotation_end = next_state%orientation(body_index)%to_matrix()
                    rotation_midpoint = &
                        force_state%orientation(body_index)%to_matrix()
                    force_state%angular_velocity(:,body_index) = &
                        matmul(transpose(rotation_midpoint), 0.5d0 * (&
                        matmul(rotation_start, state%angular_velocity(:,body_index)) + &
                        matmul(rotation_end, next_state%angular_velocity(:,body_index))))
                end do
                call force_function(force_state%time, force_state, &
                    applied_force, applied_torque, args)
            end select
        end if

        ! Assemble each body's translational and rotational discrete momentum
        ! balance. The rotational expression is equation (19) of the paper.
        do body_index = 1, nbody
            first = 6 * body_index - 5
            value(first:first+2) = bodies(body_index)%mass * &
                (x(first:first+2) - state%velocity(:,body_index)) / dt - &
                applied_force(:,body_index)

            omega1 = state%angular_velocity(:,body_index)
            omega2 = x(first+3:first+5)
            momentum1 = matmul(bodies(body_index)%inertia, omega1)
            momentum2 = matmul(bodies(body_index)%inertia, omega2)
            value(first+3:first+5) = cross_product(omega2, momentum2) + &
                sqrt(4.0d0 / dt**2 - dot_product(omega2, omega2)) * &
                momentum2 - cross_product(omega1, momentum1) - &
                sqrt(4.0d0 / dt**2 - dot_product(omega1, omega1)) * &
                momentum1 - 2.0d0 * applied_torque(:,body_index)
        end do

            ! Apply constraint forces through the reduced configuration Jacobian
            ! and append the position-level constraints at the new configuration.
        if (nconstraint > 0) then
            call constraint(next_state, constraint_value, args)
            value(1:6*nbody) = value(1:6*nbody) - matmul( &
                transpose(constraint_gradient), x(6*nbody+1:nvar))
            value(6*nbody+1:nvar) = constraint_value
        end if
    end subroutine

    subroutine state_from_unknown(x, next_state)
        !! Constructs a maximal-coordinate state from a Newton unknown vector.
        real(real64), intent(in), dimension(:) :: x
            !! The trial velocity and multiplier vector.
        type(variational_state), intent(out) :: next_state
            !! The state implied by the trial body velocities.

        ! Local Variables
        integer(int32) :: body_index, first
        type(quaternion) :: increment

        ! Apply the translational update and the unit quaternion retraction to
        ! each body independently.
        next_state = state
		next_state%time = state%time + dt
        do body_index = 1, nbody
            first = 6 * body_index - 5
            next_state%velocity(:,body_index) = x(first:first+2)
            next_state%angular_velocity(:,body_index) = x(first+3:first+5)
            next_state%position(:,body_index) = state%position(:,body_index) + &
                dt * next_state%velocity(:,body_index)
            increment = quaternion_increment( &
                next_state%angular_velocity(:,body_index), dt)
            next_state%orientation(body_index) = &
                state%orientation(body_index) * increment
            call next_state%orientation(body_index)%normalize()
        end do
    end subroutine

    subroutine finite_difference_jacobian(x, value, derivative)
        !! Computes the full Newton Jacobian by scaled forward differences.
        real(real64), intent(in), dimension(:) :: x, value
            !! The current unknown vector and its residual.
        real(real64), intent(out), dimension(:,:) :: derivative
            !! The numerical Newton Jacobian.

        ! Local Variables
        integer(int32) :: column
        real(real64) :: step_size
        ! Perturb each unknown by a scale-aware step with an absolute floor.
        do column = 1, size(x)
            step_size = this%settings%finite_difference_step * &
                max(1.0d0, abs(x(column)))
            perturbed = x
            perturbed(column) = perturbed(column) + step_size
            call evaluate_residual(perturbed, perturbed_value)
            derivative(:,column) = (perturbed_value - value) / step_size
        end do
    end subroutine

    subroutine finite_difference_constraint_gradient(derivative)
        !! Computes the reduced constraint Jacobian by perturbing translations
        !! directly and orientations on the unit-quaternion manifold.
        real(real64), intent(out), dimension(:,:) :: derivative
            !! The nconstraint-by-(6*nbody) reduced Jacobian.

        ! Local Variables
        type(variational_state) :: perturbed_state
        type(quaternion) :: perturbation
        integer(int32) :: body_index, component, column
        real(real64) :: rotation_step, translation_step
        ! Evaluate the unperturbed constraint once for all forward differences.
        call constraint(state, base_constraint_value, args)
		rotation_step = this%settings%finite_difference_step * &
			this%settings%constraint_rotation_scale
        do body_index = 1, nbody
            do component = 1, 3
                ! Translational tangent direction.
				translation_step = this%settings%finite_difference_step * &
					max(abs(state%position(component,body_index)), &
					this%settings%constraint_translation_scale)
                column = 6 * body_index - 6 + component
                perturbed_state = state
                perturbed_state%position(component,body_index) = &
                    perturbed_state%position(component,body_index) + &
					translation_step
                call constraint(perturbed_state, &
                    perturbed_value(1:nconstraint), args)
                derivative(:,column) = (perturbed_value(1:nconstraint) - &
                    base_constraint_value) / &
					translation_step

                ! Local quaternion-vector tangent direction. The scalar part
                ! maintains a unit quaternion without a separate normalization.
                column = 6 * body_index - 3 + component
                perturbed_state = state
                perturbation = quaternion([sqrt(1.0d0 - rotation_step**2), &
                    merge(rotation_step, 0.0d0, component == 1), &
                    merge(rotation_step, 0.0d0, component == 2), &
                    merge(rotation_step, 0.0d0, component == 3)])
                perturbed_state%orientation(body_index) = &
                    state%orientation(body_index) * perturbation
                call constraint(perturbed_state, &
                    perturbed_value(1:nconstraint), args)
                derivative(:,column) = (perturbed_value(1:nconstraint) - &
                    base_constraint_value) / &
					rotation_step
            end do
        end do
    end subroutine
end subroutine

! ------------------------------------------------------------------------------
function vi_solve(this, bodies, state, dt, ntime, constraint_count, &
    constraint, force_function, constraint_jacobian, multipliers, args, info) &
    result(rst)
    !! Computes the solution for the rigid-body system for the specified number
    !! of sequential time steps.
    class(variational_integrator), intent(in) :: this
        !! The variational integrator.
    type(rigid_body), intent(in), dimension(:) :: bodies
        !! The body mass and body-frame inertia properties.
    type(variational_state), intent(inout) :: state
        !! On input, the initial state; on output, the final completed state.
    real(real64), intent(in) :: dt
        !! The positive fixed time step.
    integer(int32), intent(in) :: ntime
        !! The number of times steps to take.
    integer(int32), intent(in), optional :: constraint_count
        !! The number of scalar equality constraints. The default is zero.
    procedure(variational_constraint), pointer, intent(in), optional :: constraint
        !! The holonomic equality-constraint callback. It is required when
        !! constraint_count is greater than zero.
    procedure(variational_force), pointer, intent(in), optional :: force_function
        !! The external force and torque callback. When omitted, all applied
        !! forces and torques are zero.
    procedure(variational_constraint_jacobian), pointer, intent(in), optional :: constraint_jacobian
        !! An optional analytic reduced constraint Jacobian. When omitted, the
        !! Jacobian is evaluated by finite differences.
    real(real64), allocatable, intent(out), optional, dimension(:,:) :: multipliers
        !! The Nconstraint-by-Ntime Lagrange multiplier history on success. On
        !! recoverable failure, contains only multipliers for completed steps;
        !! the final column requires a successful noncommitting look-ahead.
        !! A multiplier for the interval beginning at time point i is in column
        !! i because the discrete force uses the constraint Jacobian at that
        !! point.
    class(*), intent(inout), optional :: args
        !! Optional user-supplied data forwarded to all callbacks.
    type(variational_integrator_info), intent(out), optional :: info
        !! Optional diagnostics; iterations are summed over attempted steps,
        !! including the look-ahead when multipliers are requested. On failure,
        !! rst contains the initial state and all completed steps.
    type(variational_state), allocatable, dimension(:) :: rst
        !! The solution at each time step.

    ! Local Variables
    integer(int32) :: i, nconstraint
    real(real64), allocatable, dimension(:) :: mult
    real(real64), allocatable, dimension(:,:) :: multiplier_history
    type(variational_state), allocatable, dimension(:) :: trajectory_prefix
    type(variational_state) :: new_state, look_ahead_state
    type(variational_integrator_info) :: step_info

    ! Input Checking
    if (ntime < 1) error stop DYN_INVALID_INPUT_ERROR
	nconstraint = 0
	if (present(constraint_count)) nconstraint = constraint_count

    ! Process
    allocate(rst(ntime))
    rst(1) = state  ! initial state
    new_state = state
	if (present(info)) then
        info%converged = .true.
        info%iterations = 0
        info%jacobian_singular = .false.
	end if
	if (present(multipliers)) then
        allocate(multiplier_history(nconstraint, ntime))
	end if
    do i = 2, ntime
        if (present(info)) then
            call this%step(bodies, new_state, dt, &
                constraint_count = constraint_count, &
                constraint = constraint, &
                force_function = force_function, &
                constraint_jacobian = constraint_jacobian, &
                multipliers = mult, args = args, info = step_info)
            info%iterations = info%iterations + step_info%iterations
            info%jacobian_singular = info%jacobian_singular .or. &
                step_info%jacobian_singular
            if (.not.step_info%converged) then
                info%converged = .false.
                trajectory_prefix = rst(1:i-1)
                call move_alloc(trajectory_prefix, rst)
                state = new_state
                if (present(multipliers)) then
                    multipliers = multiplier_history(:,1:i-2)
                end if
                return
            end if
        else
            call this%step(bodies, new_state, dt, &
                constraint_count = constraint_count, &
                constraint = constraint, &
                force_function = force_function, &
                constraint_jacobian = constraint_jacobian, &
                multipliers = mult, args = args)
        end if
        rst(i) = new_state
        if (present(multipliers)) then
            multiplier_history(:,i - 1) = mult
        end if
        deallocate(mult)
    end do

    ! The final state has no outgoing interval in the requested solution. Take
    ! one additional step on a copy to determine its discrete multiplier, then
    ! discard the predicted state.
    if (present(multipliers)) then
        look_ahead_state = new_state
        if (present(info)) then
            call this%step(bodies, look_ahead_state, dt, &
                constraint_count = constraint_count, &
                constraint = constraint, &
                force_function = force_function, &
                constraint_jacobian = constraint_jacobian, &
                multipliers = mult, args = args, info = step_info)
            info%iterations = info%iterations + step_info%iterations
            info%jacobian_singular = info%jacobian_singular .or. &
                step_info%jacobian_singular
            if (.not.step_info%converged) then
                info%converged = .false.
                multipliers = multiplier_history(:,1:ntime-1)
                state = new_state
                return
            end if
        else
            call this%step(bodies, look_ahead_state, dt, &
                constraint_count = constraint_count, &
                constraint = constraint, &
                force_function = force_function, &
                constraint_jacobian = constraint_jacobian, &
                multipliers = mult, args = args)
        end if
        multiplier_history(:,ntime) = mult
    end if
    if (present(multipliers)) multipliers = multiplier_history
    state = new_state
end function

! ------------------------------------------------------------------------------
subroutine check_inputs(settings, bodies, state, dt, nconstraint, has_constraint)
    !! Validates dimensions, solver settings, body properties, and the
    !! quaternion angular-velocity domain required by the discrete update.
    type(variational_integrator_settings), intent(in) :: settings
        !! The numerical settings to validate.
    type(rigid_body), intent(in), dimension(:) :: bodies
        !! The rigid-body properties to validate.
    type(variational_state), intent(in) :: state
        !! The maximal-coordinate state to validate.
    real(real64), intent(in) :: dt
        !! The requested time step.
    integer(int32), intent(in) :: nconstraint
        !! The number of scalar equality constraints.
    logical, intent(in) :: has_constraint
        !! True when a constraint callback was supplied.

    ! Local Variables
    integer(int32) :: i, nbody

    ! Validate scalar controls and callback consistency.
    nbody = size(bodies)
    if (nbody < 1 .or. dt <= 0.0d0 .or. nconstraint < 0) &
        error stop DYN_INVALID_INPUT_ERROR
    if (nconstraint > 0 .and. .not.has_constraint) &
        error stop DYN_INVALID_INPUT_ERROR
    ! Validate state allocation and dimensions.
    if (.not.allocated(state%position) .or. &
        .not.allocated(state%orientation) .or. &
        .not.allocated(state%velocity) .or. &
        .not.allocated(state%angular_velocity)) error stop DYN_INVALID_INPUT_ERROR
    if (any(shape(state%position) /= [3, nbody]) .or. &
        any(shape(state%velocity) /= [3, nbody]) .or. &
        any(shape(state%angular_velocity) /= [3, nbody]) .or. &
        size(state%orientation) /= nbody) error stop DYN_ARRAY_SIZE_ERROR
    ! Validate nonlinear and linear solver settings.
    if (settings%tolerance <= 0.0d0 .or. &
        settings%finite_difference_step <= 0.0d0 .or. &
		settings%constraint_translation_scale <= 0.0d0 .or. &
		settings%constraint_rotation_scale <= 0.0d0 .or. &
        settings%finite_difference_step * &
            settings%constraint_rotation_scale >= 1.0d0 .or. &
        settings%maximum_iterations < 1 .or. &
        settings%maximum_line_search_iterations < 0) &
        error stop DYN_INVALID_INPUT_ERROR
    if (settings%linear_solver /= VI_DENSE_SOLVER .and. &
        settings%linear_solver /= VI_GRAPH_FACTORIZED_SOLVER) &
        error stop DYN_INVALID_INPUT_ERROR
    if (settings%force_evaluation /= VI_FORCE_LEFT_ENDPOINT .and. &
        settings%force_evaluation /= VI_FORCE_IMPLICIT_ENDPOINT .and. &
        settings%force_evaluation /= VI_FORCE_MIDPOINT) &
        error stop DYN_INVALID_INPUT_ERROR
    ! Validate body properties and the quaternion retraction domain
    ! |omega| < 2 / dt.
    do i = 1, nbody
        if (bodies(i)%mass <= 0.0d0) error stop DYN_INVALID_INPUT_ERROR
        if (abs(abs(state%orientation(i)) - 1.0d0) > &
            100.0d0 * epsilon(1.0d0)) error stop DYN_INVALID_INPUT_ERROR
        if (dot_product(state%angular_velocity(:,i), &
            state%angular_velocity(:,i)) >= 4.0d0 / dt**2) &
            error stop DYN_INVALID_INPUT_ERROR
    end do
end subroutine

! ------------------------------------------------------------------------------
function graph_factorized_solve(matrix, vector, nbody, nconstraint, &
    singular) result(rst)
    !! Solves a Newton system by graph-ordered block LDU elimination. Each rigid
    !! body is represented by a six-variable graph node and each scalar
    !! constraint by a one-variable node. Schur-complement updates add the fill
    !! edges required by Algorithm 2 of Brudigam et al. (2023).
    real(real64), intent(in), dimension(:,:) :: matrix
        !! The square Newton matrix.
    real(real64), intent(in), dimension(:) :: vector
        !! The right-hand-side vector.
    integer(int32), intent(in) :: nbody
        !! The number of six-variable rigid-body nodes.
    integer(int32), intent(in) :: nconstraint
        !! The number of scalar constraint nodes.
    logical, intent(out) :: singular
        !! True when no nonsingular graph block can be eliminated.
    real(real64), allocatable, dimension(:) :: rst
        !! The solution vector.

    ! Local Variables
    logical, allocatable, dimension(:) :: active
    integer(int32), allocatable, dimension(:) :: elimination_order
    integer(int32) :: candidate, candidate_degree, degree, i, i1, i2, &
        j, j1, j2, k, k1, k2, nnode, position
    real(real64), allocatable, dimension(:,:) :: work, diagonal, &
        normalized_upper
    real(real64), allocatable, dimension(:) :: reduced_rhs

    ! Initialization
    nnode = nbody + nconstraint
    allocate(work(size(matrix,1), size(matrix,2)), &
        reduced_rhs(size(vector)), rst(size(vector)), active(nnode), &
        elimination_order(nnode))
    work = matrix
    reduced_rhs = vector
    rst = 0.0d0
    active = .true.
    singular = .false.

    ! Eliminate the lowest-degree node whose current diagonal block is
    ! nonsingular. Constraint nodes become eligible after neighboring body
    ! elimination forms their Schur-complement diagonal.
    do position = 1, nnode
        candidate = 0
        candidate_degree = huge(1_int32)
        do i = 1, nnode
            if (.not.active(i)) cycle
            call graph_node_range(i, nbody, i1, i2)
            if (.not.block_is_nonsingular(work(i1:i2,i1:i2))) cycle
            degree = graph_node_degree(work, active, i, nbody)
            if (degree < candidate_degree) then
                candidate = i
                candidate_degree = degree
            end if
        end do
        if (candidate == 0) then
            singular = .true.
            return
        end if

        elimination_order(position) = candidate
        call graph_node_range(candidate, nbody, i1, i2)
        diagonal = work(i1:i2,i1:i2)
        reduced_rhs(i1:i2) = solve_block(diagonal, reduced_rhs(i1:i2))

        ! Premultiply each upper edge by the inverse diagonal block. These
        ! normalized edges are retained for the backward substitution.
        do k = 1, nnode
            if (.not.active(k) .or. k == candidate) cycle
            call graph_node_range(k, nbody, k1, k2)
            if (.not.blocks_connected(work, candidate, k, nbody)) cycle
            normalized_upper = solve_block_matrix(diagonal, work(i1:i2,k1:k2))
            work(i1:i2,k1:k2) = normalized_upper
        end do

        ! Update neighboring residuals and edge blocks. A previously absent
        ! block created here is a fill edge in the elimination graph.
        do j = 1, nnode
            if (.not.active(j) .or. j == candidate) cycle
            call graph_node_range(j, nbody, j1, j2)
            if (maxval(abs(work(j1:j2,i1:i2))) <= graph_zero_tolerance(work)) cycle
            reduced_rhs(j1:j2) = reduced_rhs(j1:j2) - &
                matmul(work(j1:j2,i1:i2), reduced_rhs(i1:i2))
            do k = 1, nnode
                if (.not.active(k) .or. k == candidate) cycle
                call graph_node_range(k, nbody, k1, k2)
                if (maxval(abs(work(i1:i2,k1:k2))) <= &
                    graph_zero_tolerance(work)) cycle
                work(j1:j2,k1:k2) = work(j1:j2,k1:k2) - &
                    matmul(work(j1:j2,i1:i2), work(i1:i2,k1:k2))
            end do
        end do
        active(candidate) = .false.
    end do

    ! Solve the normalized upper-triangular block system in reverse graph
    ! elimination order.
    do position = nnode, 1, -1
        i = elimination_order(position)
        call graph_node_range(i, nbody, i1, i2)
        rst(i1:i2) = reduced_rhs(i1:i2)
        do k = position + 1, nnode
            j = elimination_order(k)
            call graph_node_range(j, nbody, j1, j2)
            rst(i1:i2) = rst(i1:i2) - &
                matmul(work(i1:i2,j1:j2), rst(j1:j2))
        end do
    end do
end function

! ------------------------------------------------------------------------------
pure subroutine graph_node_range(node, nbody, first, last)
    !! Gets the global matrix indices occupied by a graph node.
    integer(int32), intent(in) :: node
        !! The graph node index.
    integer(int32), intent(in) :: nbody
        !! The number of rigid-body nodes.
    integer(int32), intent(out) :: first
        !! The first global matrix index in the node.
    integer(int32), intent(out) :: last
        !! The last global matrix index in the node.

    if (node <= nbody) then
        first = 6 * node - 5
        last = 6 * node
    else
        first = 6 * nbody + node - nbody
        last = first
    end if
end subroutine

! ------------------------------------------------------------------------------
function graph_node_degree(matrix, active, node, nbody) result(rst)
    !! Counts the active graph edges incident on a node.
    real(real64), intent(in), dimension(:,:) :: matrix
        !! The current Schur-complement matrix.
    logical, intent(in), dimension(:) :: active
        !! Flags identifying nodes that have not yet been eliminated.
    integer(int32), intent(in) :: node
        !! The node whose degree is requested.
    integer(int32), intent(in) :: nbody
        !! The number of rigid-body nodes.
    integer(int32) :: rst
        !! The active degree of the node.

    ! Local Variables
    integer(int32) :: i

    ! Count active numerical block connections.
    rst = 0
    do i = 1, size(active)
        if (active(i) .and. i /= node) then
            if (blocks_connected(matrix, node, i, nbody)) rst = rst + 1
        end if
    end do
end function

! ------------------------------------------------------------------------------
function blocks_connected(matrix, node1, node2, nbody) result(rst)
    !! Determines whether either directed block between two nodes is nonzero.
    real(real64), intent(in), dimension(:,:) :: matrix
        !! The current Schur-complement matrix.
    integer(int32), intent(in) :: node1
        !! The first node index.
    integer(int32), intent(in) :: node2
        !! The second node index.
    integer(int32), intent(in) :: nbody
        !! The number of rigid-body nodes.
    logical :: rst
        !! True when an edge exists between the nodes.

    ! Local Variables
    integer(int32) :: i1, i2, j1, j2
    real(real64) :: tolerance

    ! An edge is present when either directed Jacobian block is nonzero.
    call graph_node_range(node1, nbody, i1, i2)
    call graph_node_range(node2, nbody, j1, j2)
    tolerance = graph_zero_tolerance(matrix)
    rst = maxval(abs(matrix(i1:i2,j1:j2))) > tolerance .or. &
        maxval(abs(matrix(j1:j2,i1:i2))) > tolerance
end function

! ------------------------------------------------------------------------------
pure function graph_zero_tolerance(matrix) result(rst)
    !! Computes a scale-aware tolerance for detecting absent graph edges.
    real(real64), intent(in), dimension(:,:) :: matrix
        !! The matrix whose numerical scale is used.
    real(real64) :: rst
        !! The edge detection tolerance.

    rst = 100.0d0 * epsilon(1.0d0) * &
        max(1.0d0, maxval(abs(matrix)))
end function

! ------------------------------------------------------------------------------
function block_is_nonsingular(matrix) result(rst)
    !! Tests a candidate diagonal block using partial-pivot elimination.
    real(real64), intent(in), dimension(:,:) :: matrix
        !! The square candidate diagonal block.
    logical :: rst
        !! True when every elimination pivot is numerically nonzero.

    ! Local Variables
    integer(int32) :: i, pivot_index
    real(real64) :: tolerance
    real(real64), allocatable, dimension(:,:) :: work
    real(real64), allocatable, dimension(:) :: row

    ! Perform a small partial-pivot factorization solely to classify the block.
    work = matrix
    tolerance = graph_zero_tolerance(matrix)
    rst = .false.
    do i = 1, size(matrix,1)
        pivot_index = i - 1 + maxloc(abs(work(i:,i)), dim = 1)
        if (abs(work(pivot_index,i)) <= tolerance) return
        if (pivot_index /= i) then
            row = work(i,:)
            work(i,:) = work(pivot_index,:)
            work(pivot_index,:) = row
        end if
        if (i < size(matrix,1)) then
            work(i+1:,i) = work(i+1:,i) / work(i,i)
            work(i+1:,i+1:) = work(i+1:,i+1:) - &
                matmul(reshape(work(i+1:,i), [size(matrix,1)-i, 1]), &
                reshape(work(i,i+1:), [1, size(matrix,1)-i]))
        end if
    end do
    rst = .true.
end function

! ------------------------------------------------------------------------------
function solve_block(matrix, vector) result(rst)
    !! Solves a square diagonal-block system.
    real(real64), intent(in), dimension(:,:) :: matrix
        !! The square block matrix.
    real(real64), intent(in), dimension(:) :: vector
        !! The block right-hand side.
    real(real64), allocatable, dimension(:) :: rst
        !! The block solution.

    ! Local Variables
    integer(int32), allocatable, dimension(:) :: pivot
    real(real64), allocatable, dimension(:,:) :: lu

    ! Factor and solve the diagonal block.
    call lu_factor(matrix, ipvt = pivot, lu = lu)
    rst = solve_lu(lu, pivot, vector)
end function

! ------------------------------------------------------------------------------
function solve_block_matrix(matrix, right_hand_side) result(rst)
    !! Solves a square block against each column of a matrix.
    real(real64), intent(in), dimension(:,:) :: matrix
        !! The square block matrix.
    real(real64), intent(in), dimension(:,:) :: right_hand_side
        !! The block right-hand-side matrix.
    real(real64), allocatable, dimension(:,:) :: rst
        !! The block solution matrix.

    ! Local Variables
    integer(int32) :: i
    integer(int32), allocatable, dimension(:) :: pivot
    real(real64), allocatable, dimension(:,:) :: lu

    ! Reuse one block factorization for every right-hand-side column.
    allocate(rst(size(right_hand_side,1), size(right_hand_side,2)))
    call lu_factor(matrix, ipvt = pivot, lu = lu)
    do i = 1, size(right_hand_side,2)
        rst(:,i) = solve_lu(lu, pivot, right_hand_side(:,i))
    end do
end function

! ------------------------------------------------------------------------------
pure function quaternion_increment(omega, dt) result(increment)
    !! Computes the unit quaternion increment associated with a body-frame
    !! angular velocity over one variational time step:
    !!
    !! $$\Delta q=\left(\sqrt{1-\|h\omega/2\|^2},h\omega/2\right).$$
    real(real64), intent(in), dimension(3) :: omega
    real(real64), intent(in) :: dt
        !! The body-frame angular velocity and positive time step.
    type(quaternion) :: increment
        !! The unit orientation increment.

    ! Local Variables
    real(real64), dimension(3) :: vector_part
    real(real64) :: scalar_part

    ! Construct the vector part and select the positive scalar branch.
    vector_part = 0.5d0 * dt * omega
    scalar_part = sqrt(max(0.0d0, 1.0d0 - &
        dot_product(vector_part, vector_part)))
    increment = quaternion([scalar_part, vector_part])
end function

! ------------------------------------------------------------------------------
end module