/* VectorNav HSI Calibration Library v1.0.0.2
 *
 * Copyright (c) 2018 VectorNav Technologies, LLC
 *
 * This source code is proprietary to and copyrighted by VectorNav Technologies, LLC and is
 * available solely under license from VectorNav. All rights reserved. If you do not have a
 * license to use this source code from VectorNav, any use hereof is strictly prohibited by
 * U.S. law and other applicable law. This source code is restricted for use solely with the
 * products of VectorNav and as otherwise set forth in any agreement with VectorNav. Any
 * disclosure of this source code to the public or to any third party is also prohibited.
 */
#include "vnamath.h"

#if defined(_MSC_VER)
  #pragma warning(push)
  #pragma warning(disable:4255)
#endif
#include <stdlib.h>
#if defined(_MSC_VER)
  #pragma warning(pop)
#endif

#include <assert.h>
#include <math.h>
#include <string.h>

#include "vncompiler.h"

float
vn_sign(
    float v)
{
  if (v > 0.f) return 1.f;
  if (v < 0.f) return -1.f;

  return 0.f;
}

void
vn_init_v2(
    float* r,
    float c0,
    float c1)
{
  r[0] = c0;
  r[1] = c1;
}

void
vn_init_v3(
    float* r,
    float c0,
    float c1,
    float c2)
{
  r[0] = c0;
  r[1] = c1;
  r[2] = c2;
}

void
vn_init_v4(
    float* r,
    float c0,
    float c1,
    float c2,
    float c3)
{
  r[0] = c0;
  r[1] = c1;
  r[2] = c2;
  r[3] = c3;
}

void
vn_init_m2(
    float* r,
    float e00,
    float e01,
    float e10,
    float e11)
{
  #ifdef VN_COLUMN_ORDER
  exit(1);  /* Not Implemented. */
  #else
  r[0] = e00;
  r[1] = e10;
  r[2] = e01;
  r[3] = e11;
  #endif
}

void
vn_init_m3(
    float* r,
    float e00,
    float e01,
    float e02,
    float e10,
    float e11,
    float e12,
    float e20,
    float e21,
    float e22)
{
  #ifdef VN_COLUMN_ORDER
  exit(1);  /* Not Implemented. */
  #else
  r[0] = e00;
  r[1] = e10;
  r[2] = e20;
  r[3] = e01;
  r[4] = e11;
  r[5] = e21;
  r[6] = e02;
  r[7] = e12;
  r[8] = e22;
  #endif
}

void
vn_init_m4(
    float* r,
    float e00,
    float e01,
    float e02,
    float e03,
    float e10,
    float e11,
    float e12,
    float e13,
    float e20,
    float e21,
    float e22,
    float e23,
    float e30,
    float e31,
    float e32,
    float e33)
{
  #ifdef VN_COLUMN_ORDER
  exit(1);  /* Not Implemented. */
  #else
  r[0] = e00;
  r[1] = e10;
  r[2] = e20;
  r[3] = e30;
  r[4] = e01;
  r[5] = e11;
  r[6] = e21;
  r[7] = e31;
  r[8] = e02;
  r[9] = e12;
  r[10] = e22;
  r[11] = e32;
  r[12] = e03;
  r[13] = e13;
  r[14] = e23;
  r[15] = e33;
  #endif
}

void
vn_neg_m(
    const float* A,
    size_t m,
    size_t n,
    float* r)
{
  size_t i = 0;

  for (; i < m * n; i++)
    r[i] = -A[i];
}

void
vn_neg_m3(
    const float* A,
    float* r)
{
  vn_neg_m(A, 3, 3, r);
}

void
vn_mult_vs(
    const float* a,
    float s,
    size_t n,
    float* r)
{
  size_t i;

  for (i = 0; i < n; i++)
    r[i] = a[i] * s;
}

void
vn_mult_v2s(
    const float* a,
    float s,
    float* r)
{
  vn_mult_vs(a, s, 2, r);
}

void
vn_mult_v3s(
    const float* a,
    float s,
    float* r)
{
  vn_mult_vs(a, s, 3, r);
}

void
vn_mult_v4s(
    const float* a,
    float s,
    float* r)
{
  vn_mult_vs(a, s, 4, r);
}

void
vn_mult_mn(
    const float* A,
    const float* B,
    size_t m,
    size_t p,
    size_t n,
    float* r)
{
  size_t i, j, k;

  for (i = 0; i < m; i++)
    for (j = 0; j < n; j++)
      {
        r[n*i + j] = 0;

        for (k = 0; k < p; k++)
          r[n*i + j] = r[n*i + j] + A[p*i + k] * B[n*k + j];
      }
}

void
vn_mult_m2_m2(
    const float* A,
    const float*B,
    float* r)
{
  vn_mult_mn(A, B, 2, 2, 2, r);
}

void
vn_mult_m3_m3(
    const float* A,
    const float* B,
    float* r)
{
  vn_mult_mn(A, B, 3, 3, 3, r);
}

void
vn_mult_m4_m4(
    const float* A,
    const float* B,
    float* r)
{
  vn_mult_mn(A, B, 4, 4, 4, r);
}

void
vn_mult_m3_v3(
    const float* A,
    const float* b,
    float* r)
{
  vn_mult_mn(A, b, 3, 3, 1, r);
}

void
vn_add_v(
    const float* a,
    const float* b,
    size_t n,
    float* r)
{
  size_t i;

  for (i = 0; i < n; i++)
    r[i] = a[i] + b[i];
}

void
vn_add_v2(
    const float* a,
    const float* b,
    float* r)
{
  vn_add_v(a, b, 2, r);
}

void
vn_add_v3(
    const float* a,
    const float* b,
    float* r)
{
  vn_add_v(a, b, 3, r);
}

void
vn_add_v4(
    const float* a,
    const float* b,
    float* r)
{
  vn_add_v(a, b, 4, r);
}

void
vn_add_m(
    const float* A,
    const float* B,
    size_t m,
    size_t n,
    float* r)
{
  size_t i, j;

  for (i = 0; i < m; i++)
    for (j = 0; j < n; j++)
      r[n*i + j] = A[n*i + j] + B[n*i + j];
}

void
vn_add_m2(
    const float* A,
    const float* B,
    float* r)
{
  vn_add_m(A, B, 2, 2, r);
}

void
vn_add_m3(
    const float* A,
    const float* B,
    float* r)
{
  vn_add_m(A, B, 3, 3, r);
}

void
vn_add_m4(
    const float* A,
    const float* B,
    float* r)
{
  vn_add_m(A, B, 4, 4, r);
}

void
vn_set_m(
    float* A,
    size_t m,
    size_t n,
    size_t i,
    size_t j,
    float val)
{
  /* TODO: This looks funny. */

	/*#if VN_COL_ORDER*/
  VN_UNREFERENCED_PARAMETER(m);
	A[n*i + j] = val;
	/*#else
	A[m*j + i] = val;
	#endif*/
}

void
vn_set_m2(
    float* A,
    size_t i,
    size_t j,
    float val)
{
  vn_set_m(A, 2, 2, i, j, val);
}

void
vn_set_m3(
    float* A,
    size_t i,
    size_t j,
    float val)
{
  vn_set_m(A, 3, 3, i, j, val);
}

void
vn_set_m4(
    float* A,
    size_t i,
    size_t j,
    float val)
{
  vn_set_m(A, 4, 4, i, j, val);
}

float
vn_get_m(
    float* A,
    size_t m,
    size_t n,
    size_t i,
    size_t j)
{
  #ifdef VN_COLUMN_ORDER
  VN_UNREFERENCED_PARAMETER(m);
  return A[n*i + j];
  #else
  VN_UNREFERENCED_PARAMETER(n);
  return A[m*j + i];
  #endif
}

float
vn_get_m2(
    float* A,
    size_t i,
    size_t j)
{
  return vn_get_m(A, 2, 2, i, j);
}

float
vn_get_m3(
    float* A,
    size_t i,
    size_t j)
{
  return vn_get_m(A, 3, 3, i, j);
}

float
vn_get_m4(
    float* A,
    size_t i,
    size_t j)
{
  return vn_get_m(A, 4, 4, i, j);
}

void
vn_tranpose_m(
    const float* A,
    size_t m,
    size_t n,
    float* r)
{
  size_t i, j;

  for (i = 0; i < m; i++)
    for (j = 0; j < n; j++)
      r[m*j + i] = A[n*i + j];
}

void
vn_tranpose_m2(
    const float* A,
    float* r)
{
  vn_tranpose_m(A, 2, 2, r);
}

void
vn_tranpose_m3(
    const float* A,
    float* r)
{
  vn_tranpose_m(A, 3, 3, r);
}

void
vn_tranpose_m4(
    const float* A,
    float* r)
{
  vn_tranpose_m(A, 4, 4, r);
}

void
vn_scale_ms(
    const float* A,
    const float b,
    size_t m,
    size_t n,
    float* r)
{
  size_t i, j;

  for (i = 0; i < m; i++)
    for (j = 0; j < n; j++)
      r[n*i + j] = b * A[n*i + j];
}

void
vn_sub_m(
    const float* A,
    const float* B,
    size_t m,
    size_t n,
    float* r)
{
  size_t i, j;

  for (i = 0; i < m; i++)
    for (j = 0; j < n; j++)
      r[n*i + j] = A[n*i + j] - B[n*i + j];
}

void
vn_sub_m2(
    const float* A,
    const float* B,
    float* r)
{
  vn_sub_m(A, B, 2, 2, r);
}

void
vn_sub_m3(
    const float* A,
    const float* B,
    float* r)
{
  vn_sub_m(A, B, 3, 3, r);
}

void
vn_sub_m4(
    const float* A,
    const float* B,
    float* r)
{
  vn_sub_m(A, B, 4, 4, r);
}

void
vn_zero_m(
    float* A,
    size_t m,
    size_t n)
{
  size_t i = 0;

  for (; i < m * n; i++)
    A[i] = 0.f;
}

void
vn_zero_m2(
    float* A)
{
  vn_zero_m(A, 2, 2);
}

void
vn_zero_m3(
    float* A)
{
  vn_zero_m(A, 3, 3);
}

void
vn_zero_m4(
    float* A)
{
  vn_zero_m(A, 4, 4);
}

void
vn_zero_v(
    float* a,
    size_t n)
{
  size_t i;

  for (i = 0; i < n; i++)
    a[i] = 0;
}

void
vn_zero_v2(
    float* a)
{
  vn_zero_v(a, 2);
}

void
vn_zero_v3(
    float* a)
{
  vn_zero_v(a, 3);
}

void
vn_zero_v4(
    float* a)
{
  vn_zero_v(a, 4);
}

void
vn_eye_m(
    float* A,
    size_t m)
{
  vn_eye_mn(A, m, m);
}

void
vn_eye_mn(
    float* A,
    size_t m,
    size_t n)
{
  size_t i, max;

  vn_zero_m(A, m, n);

  max = m < n ? m : n;

  for (i = 0; i < max; i++)
    A[(max + 1)*i] = 1.f;
}

void
vn_eye_m2(
    float* A)
{
  vn_eye_m(A, 2);
}

void
vn_eye_m3(
    float* A)
{
  vn_eye_m(A, 3);
}

void
vn_eye_m4(
    float* A)
{
  vn_eye_m(A, 4);
}

enum VnError
vn_inverse_m(
    float* A,
    size_t n,
    float* r)
{
  size_t i, j, iPass, imx, icol, irow;
  float det, temp, pivot, factor;
  float *ac;

  /* TODO: Need to determine what to initialize 'factor' to. Just added line below to get rid of compiler warning. */
  factor = 1;

  ac = (float*) calloc(n*n, sizeof(float));
  det = 1;
  for (i = 0; i < n; i++)
    {
      for (j = 0; j < n; j++)
        {
          r[n*i + j] = 0;
          ac[n*i + j] = A[n*i + j];
        }

      r[n*i + i] = 1;
    }

  /* The current pivot row is iPass. For each pass, first find the maximum element in the pivot column. */
  for (iPass = 0; iPass < n; iPass++)
    {
      imx = iPass;
      for (irow = iPass; irow < n; irow++)
        {
          if (fabs(A[n*irow + iPass]) > fabs(A[n*imx + iPass])) imx = irow;
        }

      /* Interchange the elements of row iPass and row imx in both A and AInverse. */
      if (imx != iPass)
        {
          for (icol = 0; icol < n; icol++)
            {
              temp = r[n*iPass + icol];
              r[n*iPass + icol] = r[n*imx + icol];
              r[n*imx + icol] = temp;
              if (icol >= iPass)
                {
                  temp = A[n*iPass + icol];
                  A[n*iPass + icol] = A[n*imx + icol];
                  A[n*imx + icol] = temp;
                }
            }
        }
      /* The current pivot is now A[iPass][iPass].
      * The determinant is the product of the pivot elements. */
      pivot = A[n*iPass + iPass];
      det = det * pivot;
      if (det == 0)
        {
          memcpy((void*)A,(void*)ac,n*n*sizeof(float));
          free(ac);
          return E_ILL_CONDITIONED;
        }
      for (icol = 0; icol < n; icol++)
        {
          /* Normalize the pivot row by dividing by the pivot element. */
          r[n*iPass + icol] = r[n*iPass + icol] / pivot;
          if (icol >= iPass) A[n*iPass + icol] = A[n*iPass + icol] / pivot;
        }
      for (irow = 0; irow < n; irow++)
        /* Add a multiple of the pivot row to each row.  The multiple factor
        * is chosen so that the element of A on the pivot column is 0. */
        {
          if (irow != iPass) factor = A[n*irow + iPass];
          for (icol = 0; icol < n; icol++)
            {
              if (irow != iPass)
                {
                  r[n*irow + icol] -= factor * r[n*iPass + icol];
                  A[n*irow + icol] -= factor * A[n*iPass + icol];
                }
            }
        }
    }

  memcpy((void*)A,(void*)ac,n*n*sizeof(float));
  free(ac);

  return E_NONE;
}

enum VnError
vn_inverse_m2(
    float* A,
    float* r)
{
  return vn_inverse_m(A, 2, r);
}

enum VnError
vn_inverse_m3(
    float* A,
    float* r)
{
  return vn_inverse_m(A, 3, r);
}

enum VnError
vn_inverse_m4(
    float* A,
    float* r)
{
  return vn_inverse_m(A, 4, r);
}

float
vn_norm_v(
    float* a,
    size_t n)
{
  size_t i;
  float sum = 0.0;

  for (i = 0; i < n; i++)
    sum += a[i] * a[i];

  #if VN_HAVE_EXTRA_MATH_FUNCS
  return sqrtf(sum);
  #else
  return sqrt(sum);
  #endif
}

void
vn_normalize_v(
    float* a,
    size_t n,
    float* r)
{
  float magnitude = 1.f / vn_norm_v(a,n);
  vn_mult_vs(a, magnitude, n, r);
}

void
vn_normalize_v2(
    float* a,
    float* r)
{
  vn_normalize_v(a, 2, r);
}

void
vn_normalize_v3(
    float* a,
    float* r)
{
  vn_normalize_v(a, 3, r);
}

void
vn_normalize_v4(
    float* a,
    float* r)
{
  vn_normalize_v(a, 4, r);
}

float
vn_dot(
    float* a,
    float* b)
{
  size_t i;
  float sum = 0.0;

  for (i = 0; i < 3; i++)
    sum += a[i]*b[i];

  return sum;
}

void
vn_cross_v3(
    float* a,
    float* b,
    float* r)
{
  r[0] = a[1] * b[2] - a[2] * b[1];
  r[1] = a[2] * b[0] - a[0] * b[2];
  r[2] = a[0] * b[1] - a[1] * b[0];
}

void
vn_cross_norm_v3(
    float* a,
    float* b,
    float* r)
{
  vn_cross_v3(a, b, r);
  vn_normalize_v3(r, r);
}

void
vn_copy_v(
    const float* a,
    size_t d, float* r)
{
  size_t i;

  for (i = 0; i < d; i++)
    r[i] = a[i];
}

void
vn_copy_v2(
    const float* a,
    float* r)
{
  vn_copy_v(a, 2, r);
}

void
vn_copy_v3(
    const float* a,
    float* r)
{
  vn_copy_v(a, 3, r);
}

void
vn_copy_v4(
    const float* a,
    float* r)
{
  vn_copy_v(a, 4, r);
}

void
vn_copy_m(
    const float* a,
    size_t m,
    size_t n,
    float* r)
{
  size_t i;

  for (i = 0; i < m * n; i++)
    r[i] = a[i];
}

void
vn_copy_m2(
    const float* a,
    float* r)
{
  vn_copy_m(a, 2, 2, r);
}

void
vn_copy_m3(
    const float* a,
    float* r)
{
  vn_copy_m(a, 3, 3, r);
}

void
vn_copy_m4(
    const float* a,
    float* r)
{
  vn_copy_m(a, 4, 4, r);
}

void
vn_to_column_order_m(
    const float* A,
    size_t m,
    size_t n,
    float* r)
{
  size_t i, j;

  /* Issues if input and output are the same. */
  assert(A != r);

  for (i = 0; i < n; i++)
    {
      for (j = 0; j < m; j++)
        {
          r[m*j + i] = A[n*i + j];
        }
    }
}

void
vn_to_column_order_m2(
    const float* A,
    float* r)
{
  vn_to_column_order_m(A, 2, 2, r);
}

void
vn_to_column_order_m3(
    const float* A,
    float* r)
{
  vn_to_column_order_m(A, 3, 3, r);
}

void
vn_to_column_order_m4(
    const float* A,
    float* r)
{
  vn_to_column_order_m(A, 4, 4, r);
}

float
vn_deg2rad(float deg)
{
  return (deg * (float) VN_PI) / 180.f;
}

double
vn_deg2rad_d(
    double deg)
{
  return (deg * VN_PI) / 180.;
}

float
vn_cond(
    float *A,
    float *Ainv,
    size_t n)
{
  return vn_norm_v(A,n*n)*vn_norm_v(Ainv,n*n);
}

float
vn_std(
    float* samp,
    size_t n)
{
  float delta;
  float std = 0.0;
  float mean = 0.0;
  size_t i;

  for(i=0;i<n;i++)
    {
      delta = samp[i]-mean;
      mean += delta/(float)(i+1);
      std += delta*(samp[i]-mean);
    }

  #if VN_HAVE_EXTRA_MATH_FUNCS
  return sqrtf(std/(float)(n-1));
  #else
  return sqrt(std/(float)(n-1));
  #endif
}

float
vn_mean(
    float* samp,
    size_t n)
{
  size_t i;
  float mean = 0.0;

  for(i=0;i<n;i++)
    mean += samp[i];

  return mean/(float)n;
}

float
vn_abs_max(
    float* samp,
    size_t n)
{
  size_t i;
  float max = 0.0;
  float temp;

  for(i=0;i<n;i++)
    {
      #if VN_HAVE_EXTRA_MATH_FUNCS
      temp = fabsf(samp[i]);
      #else
      temp = fabs(samp[i]);
      #endif
      if(temp > max)
        max = temp;
    }

  return max;
}
