/** \file
 *
 *  Contains the Matrix3x3 class implementation.
 *
 *  Copyright (c) 2007,2008,2009 MBARI
 *  MBARI Proprietary Information.  All Rights Reserved
 *
 */

#include "Matrix3x3.h"

#include <math.h>

#include "utils/AuvMath.h"

/// Empty constructior
Matrix3x3::Matrix3x3( )
//: DataValue( 9 * sizeof ( double ), MATRIX3X3 )
{
    value2d_[0] = value_;
    value2d_[1] = value_ + 3;
    value2d_[2] = value_ + 6;
    *this = 0.0;
}

/// Single Value constructior
Matrix3x3::Matrix3x3( const double& init )
//: DataValue( 9 * sizeof ( double ), MATRIX3X3 )
{
    value2d_[0] = value_;
    value2d_[1] = value_ + 3;
    value2d_[2] = value_ + 6;
    *this = init;
}

/// Single Value constructior
Matrix3x3::Matrix3x3( const double value11, const double value12, const double value13,
                      const double value21, const double value22, const double value23,
                      const double value31, const double value32, const double value33 )
//: DataValue( 9 * sizeof ( double ), MATRIX3X3 )
{
    value2d_[0] = value_;
    value2d_[1] = value_ + 3;
    value2d_[2] = value_ + 6;
    value_[0] = value11;
    value_[1] = value12;
    value_[2] = value13;
    value_[3] = value21;
    value_[4] = value22;
    value_[5] = value23;
    value_[6] = value31;
    value_[7] = value32;
    value_[8] = value33;
}

/// Coordinate transformation constructor
Matrix3x3::Matrix3x3( Point6D pos )
{
    value2d_[0] = value_;
    value2d_[1] = value_ + 3;
    value2d_[2] = value_ + 6;
    float sinRol = sin( pos.getRoll() );
    float cosRol = cos( pos.getRoll() );
    float sinPch = sin( pos.getPitch() );
    float cosPch = cos( pos.getPitch() );
    float sinHdg = sin( pos.getHeading() );
    float cosHdg = cos( pos.getHeading() );
    float sinPchCosHdg = sinPch * cosHdg;
    float sinPchSinHdg = sinPch * sinHdg;

    value_[0] =  cosPch * cosHdg;
    value_[1] = -cosRol * sinHdg + sinRol * sinPchCosHdg;
    value_[2] =  sinRol * sinHdg + cosRol * sinPchCosHdg;
    value_[3] =  cosPch * sinHdg;
    value_[4] =  cosRol * cosHdg + sinRol * sinPchSinHdg;
    value_[5] = -sinRol * cosHdg + cosRol * sinPchSinHdg;
    value_[6] = -sinPch;
    value_[7] =  cosPch * sinRol;
    value_[8] =  cosPch * cosRol;

}

// as a Str
Str Matrix3x3::toString() const
{
    Str strm( "[" );
    for( int i = 0; i < 3; ++i )
    {
        if( i > 0 )
        {
            strm += ",";
        }
        strm += "[";
        for( int j = 0; j < 3; ++j )
        {
            if( j > 0 )
            {
                strm += ",";
            }
            strm += Str( value_[ i * 3 + j ] );
        }
        strm += "]";
    }
    return strm + "]";
}

void Matrix3x3::transpose()
{
    AuvMath::Swap( value_[1], value_[3] );
    AuvMath::Swap( value_[2], value_[6] );
    AuvMath::Swap( value_[5], value_[7] );
}

void Matrix3x3::invert()
{
    for( int i = 1; i < 3; i++ )        // normalize row 0
    {
        value_[ i ] /= value_[ 0 ];
    }
    for( int i = 1; i < 3; i++ )
    {
        for( int j = i; j < 3; j++ )        // do a column of L
        {
            double sum = 0.0;
            for( int k = 0; k < i; k++ )
            {
                sum += value_[ j * 3 + k ] * value_[ k * 3 + i ];
            }
            value_[ j * 3 + i ] -= sum;
        }
        if( i == 2 )
        {
            continue;
        }
        for( int j = i + 1; j < 3; j++ )        // do a row of U
        {
            double sum = 0.0;
            for( int k = 0; k < i; k++ )
            {
                sum += value_[ i * 3 + k ] * value_[ k * 3 + j ];
            }
            value_[ i * 3 + j ] = ( value_[ i * 3 + j ] - sum ) / value_[ i * 3 + i ];
        }
    }
    for( int i = 0; i < 3; i++ )        // invert L
    {
        for( int j = i; j < 3; j++ )
        {
            double x = 1.0;
            if( i != j )
            {
                x = 0.0;
                for( int k = i; k < j; k++ )
                {
                    x -= value_[ j * 3 + k ] * value_[ k * 3 + i ];
                }
            }
            value_[ j * 3 + i ] = x / value_[ j * 3 + j ];
        }
    }
    for( int i = 0; i < 3; i++ )        // invert U
    {
        for( int j = i; j < 3; j++ )
        {
            if( i == j )
            {
                continue;
            }
            double sum = 0.0;
            for( int k = i; k < j; k++ )
            {
                sum += value_[ k * 3 + j ] * ( ( i == k ) ? 1.0 : value_[ i * 3 + k ] );
            }
            value_[ i * 3 + j ] = -sum;
        }
    }
    for( int i = 0; i < 3; i++ )          // final inversion
    {
        for( int j = 0; j < 3; j++ )
        {
            double sum = 0.0;
            for( int k = ( ( i > j ) ? i : j ); k < 3; k++ )
            {
                sum += ( ( j == k ) ? 1.0 : value_[ j * 3 + k ] ) * value_[ k * 3 + i ];
            }
            value_[ j * 3 + i ] = sum;
        }
    }
};

void Matrix3x3::decomposeLU()
{

    for( int k = 0; k < 3; k++ )   //loop over columns of A
    {
        for( int i = 0; i <= k; i++ )
        {
            //loop over elements of U matrix in column k
            for( int m = 0; m < i; m++ )
            {
                value_[ i * 3 + k ] = value_[ i * 3 + k ] - value_[ i * 3 + m ] * value_[ m * 3 + k ];
            }
        }
        for( int i = k + 1; i < 3; i++ )
        {
            //loop over elements of U matrix in column k
            for( int m = 0; m < k; m++ )
            {
                value_[ i * 3 + k ] = value_[ i * 3 + k ] - value_[ i * 3 + m ] * value_[ m * 3 + k ];
            }
            value_[ i * 3 + k ] = value_[ i * 3 + k ] / value_[ k * 4 ];
        }
    }
};

// Solve for x in Ax = b, where this is the LU decompostion of
// M and b is supplied here.
void Matrix3x3::solveLU( Point3D& b, Point3D& x ) const
{
    double y[ 3 ];
    int i, j;

    /*
      solve Ly = b by forward substitution
    */
    for( i = 0; i < 3; ++i )
    {
        y[ i ] = b[ i ];
        for( j = 0; j <= i - 1; ++j )
        {
            y[ i ] -= value_[ i * 3 + j ] * y[ j ];
        }
    }
    /*
      solve Ux = y by backward substitution
    */
    for( i = 2; i >= 0; --i )
    {
        x[ i ] = y[ i ];
        for( j = 2; j >= i + 1; --j )
        {
            x[ i ] -= value_[ i * 3 + j ] * x[ j ];
        }
        x[ i ] /= value_[ i * 3 + i ];
    }

}

const void Matrix3x3::mult( Point3D & lhs, Point3D & rhs ) const
{
    double * r = rhs.getPtr1d();
    lhs.setX( value_[ 0 ] * r[ 0 ] + value_[ 1 ] * r[ 1 ] + value_[ 2 ] * r[ 2 ] );
    lhs.setY( value_[ 3 ] * r[ 0 ] + value_[ 4 ] * r[ 1 ] + value_[ 5 ] * r[ 2 ] );
    lhs.setZ( value_[ 6 ] * r[ 0 ] + value_[ 7 ] * r[ 1 ] + value_[ 8 ] * r[ 2 ] );
}

bool Matrix3x3::operator==( const Matrix3x3 & rhs )
{
    for( int i = 0; i < 9; ++i )
    {
        if( value_[ i ] != rhs.value_[ i ] )
        {
            return false;
        }
    }
    return true;
}

Matrix3x3& Matrix3x3::operator=( const Matrix3x3 & rhs )
{
    for( int i = 0; i < 9; ++i )
    {
        value_[ i ] = rhs.value_[ i ];
    }
    return *this;
}

Matrix3x3& Matrix3x3::operator=( const double rhs )
{
    for( int i = 0; i < 9; ++i )
    {
        value_[ i ] = rhs;
    }
    return *this;
}

Matrix3x3& Matrix3x3::operator+=( const Matrix3x3 & rhs )
{
    for( int i = 0; i < 9; ++i )
    {
        value_[ i ] += rhs.value_[ i ];
    }
    return *this;
}

Matrix3x3& Matrix3x3::operator-=( const Matrix3x3 & rhs )
{
    for( int i = 0; i < 9; ++i )
    {
        value_[ i ] -= rhs.value_[ i ];
    }
    return *this;
}

Matrix3x3& Matrix3x3::operator*=( const double rhs )
{
    for( int i = 0; i < 9; ++i )
    {
        value_[ i ] *= rhs;
    }
    return *this;
}

void Matrix3x3::set( int index1, int index2, const double value )
{
    if( index1 < 0 || index1 > 2 )
    {
        return;
    }
    if( index2 < 0 || index2 > 2 )
    {
        return;
    }
    value_[ index1 * 3 + index2 ] = value;
}

double& Matrix3x3::operator()( int index1, int index2 )
{
    if( index1 < 0 || index1 > 2 )
    {
        index1 = 0;
    }
    if( index2 < 0 || index2 > 2 )
    {
        index2 = 0;
    }
    return value_[ index1 * 3 + index2 ];
}
