//===----------------------------------------------------------------------===//
//
// Part of the DirectXShaderCompiler, under the Apache License v2.0 with LLVM
// Exceptions.
// See https://llvm.org/LICENSE.txt for license information.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
//===----------------------------------------------------------------------===//
// Template trait constructs for driving template specialization.
//===----------------------------------------------------------------------===//

#ifndef __HLSL_TYPE_TRAITS
#define __HLSL_TYPE_TRAITS

#if __HLSL_VERSION >= 2021

#define SIZE_TYPE int

namespace hlsl {

template <typename T, typename U> struct is_same {
  static const bool value = false;
};

template <typename T> struct is_same<T, T> {
  static const bool value = true;
};

template <typename T> struct is_arithmetic {
  static const bool value = false;
};

#define __ARITHMETIC_TYPE(type)                                                \
  template <> struct is_arithmetic<type> {                                     \
    static const bool value = true;                                            \
  };

#if __HLSL_ENABLE_16_BIT
__ARITHMETIC_TYPE(uint16_t)
__ARITHMETIC_TYPE(int16_t)
#endif
__ARITHMETIC_TYPE(uint)
__ARITHMETIC_TYPE(int)
__ARITHMETIC_TYPE(uint64_t)
__ARITHMETIC_TYPE(int64_t)
__ARITHMETIC_TYPE(half)
__ARITHMETIC_TYPE(float)
__ARITHMETIC_TYPE(double)

#undef __ARITHMETIC_TYPE

template <typename T> struct is_signed {
  static const bool value = true;
};

#define __UNSIGNED_TYPE(type)                                                  \
  template <> struct is_signed<type> {                                         \
    static const bool value = false;                                           \
  };

#if __HLSL_ENABLE_16_BIT
__UNSIGNED_TYPE(uint16_t)
#endif
__UNSIGNED_TYPE(uint)
__UNSIGNED_TYPE(uint64_t)

#undef __UNSIGNED_TYPE

template <typename T> struct is_vector {
  static const bool value = false;
};

template <typename T, uint N> struct is_vector<vector<T, N> > {
  static const bool value = true;
};

template <typename T> struct strip_vector_type {
  using type = T;
};

template <typename T, SIZE_TYPE N> struct strip_vector_type<vector<T, N> > {
  using type = T;
};

template <typename T> struct is_arithmetic_vector {
  static const bool value =
      is_arithmetic<T>::value ||
      (is_vector<T>::value &&
       is_arithmetic<typename strip_vector_type<T>::type>::value);
};

} // namespace hlsl

#endif // __HLSL_VERSION >= 2021
#endif // __HLSL_TYPE_TRAITS
