Stratax 0.3.1
Loading...
Searching...
No Matches
Validation.hpp
1#pragma once
2
3#include <cstddef>
4#include <limits>
5#include <type_traits>
6
7#include <stratax/core/dtypes/Concepts.hpp>
8#include <stratax/exceptions/Exceptions.hpp>
9
10namespace stratax::core::validation {
11
12inline std::size_t nonnegative_size(std::ptrdiff_t value, const char* message)
13{
14 if (value < 0)
15 {
16 throw Exceptions::DimensionError(message);
17 }
18
19 return static_cast<std::size_t>(value);
20}
21
22inline void require_rank(std::size_t actual, std::size_t expected, const char* message)
23{
24 if (actual != expected)
25 {
26 throw Exceptions::DimensionError(message);
27 }
28}
29
30template<typename Ranked>
31requires requires(const Ranked& object)
32{
33 object.rank();
34}
35const Ranked& require_rank(const Ranked& object, std::size_t expected, const char* message)
36{
37 require_rank(object.rank(), expected, message);
38 return object;
39}
40
41inline std::size_t checked_multiply(
42 std::size_t lhs,
43 std::size_t rhs,
44 const char* message)
45{
46 if (rhs != 0 && lhs > std::numeric_limits<std::size_t>::max() / rhs)
47 {
48 throw Exceptions::DimensionError(message);
49 }
50
51 return lhs * rhs;
52}
53
54inline std::size_t checked_add(
55 std::size_t lhs,
56 std::size_t rhs,
57 const char* message)
58{
59 if (lhs > std::numeric_limits<std::size_t>::max() - rhs)
60 {
61 throw Exceptions::DimensionError(message);
62 }
63
64 return lhs + rhs;
65}
66
67inline std::size_t nonnegative_index(std::ptrdiff_t value, const char* message)
68{
69 if (value < 0)
70 {
71 throw Exceptions::IndexError(message);
72 }
73
74 return static_cast<std::size_t>(value);
75}
76
77inline void require_index(std::size_t index, std::size_t size, const char* message)
78{
79 if (index >= size)
80 {
81 throw Exceptions::IndexError(message);
82 }
83}
84
85inline void require_at_most(std::size_t value, std::size_t upper, const char* message)
86{
87 if (value > upper)
88 {
89 throw Exceptions::IndexError(message);
90 }
91}
92
93inline std::size_t nonnegative_shape_dimension(std::ptrdiff_t value, const char* message)
94{
95 if (value < 0)
96 {
97 throw Exceptions::ShapeError(message);
98 }
99
100 return static_cast<std::size_t>(value);
101}
102
103inline std::size_t positive_shape_dimension(std::ptrdiff_t value, const char* message)
104{
105 if (value <= 0)
106 {
107 throw Exceptions::ShapeError(message);
108 }
109
110 return static_cast<std::size_t>(value);
111}
112
113inline void require_positive_shape_dimension(std::size_t value, const char* message)
114{
115 if (value == 0)
116 {
117 throw Exceptions::ShapeError(message);
118 }
119}
120
121template<typename Lhs, typename Rhs>
122[[nodiscard]] bool same_shape(const Lhs& lhs, const Rhs& rhs)
123{
124 return lhs.size() == rhs.size() && lhs.shape() == rhs.shape();
125}
126
127template<typename Lhs, typename Rhs>
128void require_same_shape(const Lhs& lhs, const Rhs& rhs, const char* message)
129{
130 if (!same_shape(lhs, rhs))
131 {
132 throw Exceptions::ShapeError(message);
133 }
134}
135
136inline void require_equal_size(std::size_t lhs, std::size_t rhs, const char* message)
137{
138 if (lhs != rhs)
139 {
140 throw Exceptions::ShapeError(message);
141 }
142}
143
144template<typename Actual, typename Expected>
145void require_type(const char* message)
146{
147 if constexpr (!std::same_as<std::remove_cvref_t<Actual>, std::remove_cvref_t<Expected>>)
148 {
149 throw Exceptions::TypeError(message);
150 }
151}
152
153template<typename T>
154void require_numeric_type(const char* message)
155{
156 if constexpr (!Numeric<T>)
157 {
158 throw Exceptions::TypeError(message);
159 }
160}
161
162template<typename Lhs, typename Rhs>
163requires requires
164{
165 typename Lhs::value_type;
166 typename Rhs::value_type;
167}
168void require_same_value_type(const Lhs& lhs, const Rhs& rhs, const char* message)
169{
170 (void)lhs;
171 (void)rhs;
172 require_type<typename Lhs::value_type, typename Rhs::value_type>(message);
173}
174
175} // namespace stratax::core::validation