#include <c6x.h>
class int16x8_t;
class int16x4_t;
class uint8x4_t;

const int64_t   mask = 0x0101010101010101;
const uint32_t  mask_a = _loll(mask);

class uint8x4_t {
public:
    uint32_t data;
    uint8x4_t(uint32_t val);
    uint8x4_t();
    inline uint8x4_t operator= (const uint8_t input[restrict]);
    static inline uint8x4_t lower32 (int16x4_t src);
    static inline uint8x4_t upper32 (int16x4_t src);
    static inline uint8x4_t add16 (uint8x4_t src0, uint8x4_t src1);
    inline void store (uint8_t output[restrict]);
};

class int16x4_t {
public:
    int64_t data;
    inline int16x4_t(int64_t val);
    inline int16x4_t();
    
    inline int16x4_t operator= (const uint8_t input[restrict]); 
    static inline int16x4_t lower64 (int16x8_t src);
    static inline int16x4_t upper64 (int16x8_t src);
    static inline int16x4_t add16 (int16x4_t src0, int16x4_t src1);
    static inline int16x4_t multiply_8bitwise(uint8x4_t src0, uint8x4_t src1);
    inline void store (uint8_t output[restrict]);
};

class int16x8_t {
public:
    __x128_t data;
    inline int16x8_t(__x128_t val);
    
    inline int16x8_t operator= (const uint8_t input[restrict]);
    static inline int16x8_t multiply_8bitwise(const int16x4_t src0, const int16x4_t src1);
};

inline uint8x4_t::uint8x4_t(uint32_t val) {
    data = val;
} 

inline uint8x4_t uint8x4_t::operator= (const uint8_t input[restrict]) {
    data = _amem4_const(input);
    return uint8x4_t(data);
}

inline uint8x4_t uint8x4_t::lower32 (int16x4_t src) {
    return uint8x4_t(_loll(src.data));
}

inline uint8x4_t uint8x4_t::upper32 (int16x4_t src) {
    return uint8x4_t(_hill(src.data));
}

inline uint8x4_t uint8x4_t::add16 (uint8x4_t src0, uint8x4_t src1) {
    return uint8x4_t(_add2((int32_t)src0.data, (int32_t)src1.data));
}


inline int16x4_t::int16x4_t(int64_t val) : data(val) {}
inline int16x4_t::int16x4_t() {}

inline int16x4_t int16x4_t::operator= (const uint8_t input[restrict]) {
    data = _mem8_const(input);
    return int16x4_t(data);
}    

inline int16x4_t int16x4_t::lower64 (int16x8_t src) {
  return int16x4_t(_lo128(src.data));
}

inline int16x4_t int16x4_t::upper64 (int16x8_t src) {
  return int16x4_t(_hi128(src.data));
}

inline int16x4_t int16x4_t::add16 (int16x4_t src0, int16x4_t src1) {
    return int16x4_t(_dadd2(src0.data, src1.data));
}

inline int16x4_t int16x4_t::multiply_8bitwise(uint8x4_t src0, uint8x4_t src1) {
    return int16x4_t(_mpyu4ll(src0.data, src1.data));
}

inline void int16x4_t::store (uint8_t output[restrict]) {
    _mem8((void *)output) = data;
}

inline int16x8_t::int16x8_t(__x128_t val) : data(val) {}

inline int16x8_t int16x8_t::operator= (const uint8_t input[restrict]) {
    data = _llto128(_amem8_const(input+8),_amem8_const(input));
    return int16x8_t(data);
}

inline int16x8_t int16x8_t::multiply_8bitwise(const int16x4_t src0, const int16x4_t src1) {
    return int16x8_t(_dmpyu4((int64_t)src0.data, src1.data));
}

namespace pixellib_filter {
  
  class box3x3 {
  public:
      uint8x4_t xaxb;
      inline box3x3();
      inline box3x3(uint8x4_t row0, uint8x4_t row1, uint8x4_t row2);
      inline void filter(int16x4_t row0, int16x4_t row1, 
                         int16x4_t row2, int16x4_t &dst);
      
  };
}

inline pixellib_filter::box3x3::box3x3() {}

inline void pixellib_filter::box3x3::filter(int16x4_t row0, int16x4_t row1, int16x4_t row2,
                                            int16x4_t &dst) {


    /* Convert from 8bpp to 16bpp so we can do SIMD addition of rows */
    int16x8_t r0_2 = int16x8_t::multiply_8bitwise (row0, int16x4_t(mask));

    int16x4_t test = int16x4_t::lower64(r0_2);

    dst = test;
}

void vxBoxKernel(const uint8_t * restrict src,
                       uint8_t * restrict dst0,
                       int32_t srcStride,
                       int32_t outWidth) {
    
    int16x4_t src_8a, src_8b, src_8c;
    int16x4_t dst_8a;    

    pixellib_filter::box3x3 box0;

//#pragma UNROLL(2)
//#pragma MUST_ITERATE(0,,8)
    for (int x=0;x<1024;x++) {
        /* Read 8 bytes from each of the 3 lines. */
        src_8a = &src[(srcStride*0) + (x*8) + 2];
        src_8b = &src[(srcStride*1) + (x*8) + 2];
        src_8c = &src[(srcStride*2) + (x*8) + 2];

        box0.filter(src_8a, src_8b, src_8c, dst_8a);

        dst_8a.store(&dst0[x*8]);
    }
}
