#include "crMath.h"

namespace crMath
{
    float4x4 Mult4x4( const float4x4& a, const float4x4& b )
    {
        float4x4 r = {};
        for( int32_t c = 0; c < 4; ++c )
        {
            for( int32_t row = 0; row < 4; ++row )
            {
                float sum = 0.0f;
                for( int32_t k = 0; k < 4; ++k )
                    sum += ( a.m[ ( k * 4 ) + row ] * b.m[ ( c * 4 ) + k ] );

                r.m[ ( c * 4 ) + row ] = sum;
            }
        }
        return r;
    }

    float4x4 OrthoRH( float halfWidth, float halfHeight, float nearPlane, float farPlane )
    {
        float4x4 r = {};
        r.m[ 0 ]  = 1.0f / halfWidth;
        r.m[ 5 ]  = 1.0f / halfHeight;
        r.m[ 10 ] = -2.0f / ( farPlane - nearPlane );
        r.m[ 14 ] = -( farPlane + nearPlane ) / ( farPlane - nearPlane );
        r.m[ 15 ] = 1.0f;
        return r;
    }
    float4x4 LookAtRH( float3 eye, float3 target, float3 up )
    {
        const float3 fwd = Normalize3( float3( target.x - eye.x, target.y - eye.y, target.z - eye.z ) );
        const float3 s   = Normalize3( Cross3( fwd, up ) );
        const float3 u   = Cross3( s, fwd );

        float4x4 r = {};
        r.m[ 0 ]  = s.x;
        r.m[ 4 ]  = s.y;
        r.m[ 8 ]  = s.z;
        r.m[ 1 ]  = u.x;
        r.m[ 5 ]  = u.y;
        r.m[ 9 ]  = u.z;
        r.m[ 2 ]  = -fwd.x;
        r.m[ 6 ]  = -fwd.y;
        r.m[ 10 ] = -fwd.z;
        r.m[ 12 ] = -Dot3( s, eye );
        r.m[ 13 ] = -Dot3( u, eye );
        r.m[ 14 ] = Dot3( fwd, eye );
        r.m[ 15 ] = 1.0f;
        return r;
    }

    // column-major storage: math-row r = ( m[r], m[r+4], m[r+8], m[r+12] ); the clip-space
    // inequalities -w <= {x,y,z} <= w become (row3 +/- rowN) . p >= 0. planes are normalized so
    // the plane<->center distance is in world units and compares directly against the sphere radius
    static crPlane MakePlane( float a, float b, float c, float d )
    {
        const float len = Sqrt( ( a * a ) + ( b * b ) + ( c * c ) );
        const float inv = ( len > 0.000001f ) ? ( 1.0f / len ) : 0.0f;

        crPlane p;
        p.n = float3( a * inv, b * inv, c * inv );
        p.d = ( d * inv );
        return p;
    }
    crFrustum FrustumFromMatrix( const float4x4& vp )
    {
        const float* m = vp.m;
        const float  r0[ 4 ] = { m[ 0 ], m[ 4 ], m[  8 ], m[ 12 ] };
        const float  r1[ 4 ] = { m[ 1 ], m[ 5 ], m[  9 ], m[ 13 ] };
        const float  r2[ 4 ] = { m[ 2 ], m[ 6 ], m[ 10 ], m[ 14 ] };
        const float  r3[ 4 ] = { m[ 3 ], m[ 7 ], m[ 11 ], m[ 15 ] };

        crFrustum f;
        f.planes[ 0 ] = MakePlane( r3[ 0 ] + r0[ 0 ], r3[ 1 ] + r0[ 1 ], r3[ 2 ] + r0[ 2 ], r3[ 3 ] + r0[ 3 ] );   // left
        f.planes[ 1 ] = MakePlane( r3[ 0 ] - r0[ 0 ], r3[ 1 ] - r0[ 1 ], r3[ 2 ] - r0[ 2 ], r3[ 3 ] - r0[ 3 ] );   // right
        f.planes[ 2 ] = MakePlane( r3[ 0 ] + r1[ 0 ], r3[ 1 ] + r1[ 1 ], r3[ 2 ] + r1[ 2 ], r3[ 3 ] + r1[ 3 ] );   // bottom
        f.planes[ 3 ] = MakePlane( r3[ 0 ] - r1[ 0 ], r3[ 1 ] - r1[ 1 ], r3[ 2 ] - r1[ 2 ], r3[ 3 ] - r1[ 3 ] );   // top
        f.planes[ 4 ] = MakePlane( r3[ 0 ] + r2[ 0 ], r3[ 1 ] + r2[ 1 ], r3[ 2 ] + r2[ 2 ], r3[ 3 ] + r2[ 3 ] );   // near
        f.planes[ 5 ] = MakePlane( r3[ 0 ] - r2[ 0 ], r3[ 1 ] - r2[ 1 ], r3[ 2 ] - r2[ 2 ], r3[ 3 ] - r2[ 3 ] );   // far
        return f;
    }

    b3Quat IntegrateRotation( b3Quat q, float3 angularVelocity, float dt )
    {
        // dq/dt = 0.5 * (w quat q) — same scheme box2d used in 2D, then renormalize
        const b3Quat wq = { { angularVelocity.x, angularVelocity.y, angularVelocity.z }, 0.0f };
        const b3Quat dq = b3MulQuat( wq, q );

        const float half = ( 0.5f * dt );

        b3Quat r;
        r.v.x = q.v.x + ( half * dq.v.x );
        r.v.y = q.v.y + ( half * dq.v.y );
        r.v.z = q.v.z + ( half * dq.v.z );
        r.s   = q.s   + ( half * dq.s );

        return b3NormalizeQuat( r );
    }

    uint64_t HashFnv1a( const void* data, size_t size, uint64_t seed )
    {
        const uint8_t* bytes = static_cast<const uint8_t*>( data );

        uint64_t hash = seed;
        for( size_t i = 0; i < size; ++i )
        {
            hash ^= bytes[ i ];
            hash *= 1099511628211ULL;   // FNV prime
        }
        return hash;
    }
}
