1use std::ops;
2
3use nalgebra::{SimdRealField, SVector};
4
5#[derive(Debug, Clone, PartialEq)]
8pub struct BVector<T: SimdRealField + Copy, const N: usize> {
9 pub data_vector: SVector<T, N>,
10}
11
12impl<T: SimdRealField + Copy, const N: usize> BVector<T, N> {
13 pub fn from_element(element: T) -> Self {
15 BVector {
16 data_vector: SVector::<T, N>::from_element(element),
17 }
18 }
19
20 pub fn from_array(data: [T; N]) -> Self
22 where
23 T: Copy,
24 {
25 BVector {
26 data_vector: SVector::<T, N>::from_row_slice(&data),
27 }
28 }
29
30 pub fn dot(&self, other: &Self) -> T {
31 self.data_vector.dot(&other.data_vector)
32 }
33
34 pub fn len(&self) -> usize {
35 N
36 }
37
38 pub fn is_empty(&self) -> bool {
40 N == 0
41 }
42
43 pub fn get(&self, index: usize) -> T {
44 assert!(index < N, "Index out of bounds");
45 self.data_vector[index]
46 }
47}
48
49impl<T: SimdRealField + Copy, const N: usize> ops::Add for BVector<T, N> {
50 type Output = BVector<T, N>;
51
52 fn add(self, rhs: Self) -> Self::Output {
53 BVector {
54 data_vector: self.data_vector + rhs.data_vector,
55 }
56 }
57}
58
59#[macro_export]
66macro_rules! bvector {
67 ($($x:expr),+ $(,)?) => {
68 {
69 $crate::BVector::from_array([$( $x ),+])
71 }
72 };
73}
74
75
76#[cfg(test)]
77mod tests {
78 use super::*;
79
80 #[test]
81 fn from_element_fills_every_position() {
82 let v = BVector::<f64, 3>::from_element(1.5);
83
84 assert_eq!((v.get(0), v.get(1), v.get(2)), (1.5, 1.5, 1.5));
85 }
86
87 #[test]
88 fn the_macro_keeps_elements_in_order() {
89 let v = bvector![1.0, 2.0, 3.0];
90
91 assert_eq!(v, BVector::from_array([1.0, 2.0, 3.0]));
92 assert_eq!((v.get(0), v.get(1), v.get(2)), (1.0, 2.0, 3.0));
93 }
94
95 #[test]
96 fn the_macro_accepts_a_trailing_comma() {
97 assert_eq!(bvector![1.0, 2.0,], bvector![1.0, 2.0]);
98 }
99
100 #[test]
101 fn len_is_the_compile_time_dimension() {
102 assert_eq!(bvector![1.0, 2.0, 3.0].len(), 3);
103 assert!(!bvector![1.0].is_empty());
104 assert!(BVector::<f64, 0>::from_array([]).is_empty());
105 }
106
107 #[test]
108 fn dot_sums_the_element_wise_products() {
109 assert_eq!(bvector![1.0, 2.0, 3.0].dot(&bvector![4.0, -5.0, 6.0]), 12.0);
110 }
111
112 #[test]
113 fn add_is_element_wise() {
114 assert_eq!(bvector![1.0, 2.0] + bvector![0.5, -3.0], bvector![1.5, -1.0]);
115 }
116
117 #[test]
118 #[should_panic(expected = "Index out of bounds")]
119 fn get_rejects_an_index_past_the_end() {
120 bvector![1.0, 2.0].get(2);
121 }
122}