Skip to main content

oxmpl_js/base/
real_vector_state_space.rs

1// Copyright (c) 2025 Ross Gardiner, Junior Sundar
2//
3// SPDX-License-Identifier: BSD-3-Clause
4
5use oxmpl::base::{
6    space::{RealVectorStateSpace as OxmplRealVectorStateSpace, StateSpace},
7    state::RealVectorState as OxmplRealVectorState,
8};
9use rand::rng;
10use std::sync::{Arc, Mutex};
11use wasm_bindgen::prelude::*;
12
13use crate::base::JsRealVectorState;
14
15#[wasm_bindgen(js_name = RealVectorStateSpace)]
16pub struct JsRealVectorStateSpace {
17    #[wasm_bindgen(skip)]
18    pub inner: Arc<Mutex<OxmplRealVectorStateSpace>>,
19}
20
21#[wasm_bindgen(js_class = RealVectorStateSpace)]
22impl JsRealVectorStateSpace {
23    #[wasm_bindgen(constructor)]
24    pub fn new(
25        dimension: usize,
26        bounds: Option<Vec<f64>>,
27    ) -> Result<JsRealVectorStateSpace, String> {
28        let bounds_vec = if let Some(b) = bounds {
29            if b.len() != dimension * 2 {
30                return Err(format!(
31                    "Bounds array must have {} elements (2 per dimension)",
32                    dimension * 2
33                ));
34            }
35            let mut bounds_tuples = Vec::new();
36            for i in 0..dimension {
37                bounds_tuples.push((b[i * 2], b[i * 2 + 1]));
38            }
39            Some(bounds_tuples)
40        } else {
41            None
42        };
43
44        match OxmplRealVectorStateSpace::new(dimension, bounds_vec) {
45            Ok(space) => Ok(Self {
46                inner: Arc::new(Mutex::new(space)),
47            }),
48            Err(e) => Err(e.to_string()),
49        }
50    }
51
52    #[wasm_bindgen(js_name = sample)]
53    pub fn sample(&self) -> Result<JsRealVectorState, String> {
54        let mut rng = rng();
55        match self.inner.lock().unwrap().sample_uniform(&mut rng) {
56            Ok(state) => Ok(JsRealVectorState::new(state.values)),
57            Err(e) => Err(e.to_string()),
58        }
59    }
60
61    #[wasm_bindgen(js_name = distance)]
62    pub fn distance(&self, state1: &JsRealVectorState, state2: &JsRealVectorState) -> f64 {
63        self.inner
64            .lock()
65            .unwrap()
66            .distance(&state1.inner, &state2.inner)
67    }
68
69    #[wasm_bindgen(js_name = satisfiesBounds)]
70    pub fn satisfies_bounds(&self, state: &JsRealVectorState) -> bool {
71        self.inner.lock().unwrap().satisfies_bounds(&state.inner)
72    }
73
74    #[wasm_bindgen(js_name = enforceBounds)]
75    pub fn enforce_bounds(&self, state: &JsRealVectorState) -> JsRealVectorState {
76        let mut new_state = (*state.inner).clone();
77        self.inner.lock().unwrap().enforce_bounds(&mut new_state);
78        JsRealVectorState {
79            inner: Arc::new(new_state),
80        }
81    }
82
83    #[wasm_bindgen(js_name = interpolate)]
84    pub fn interpolate(
85        &self,
86        from: &JsRealVectorState,
87        to: &JsRealVectorState,
88        t: f64,
89    ) -> JsRealVectorState {
90        let mut result_state = OxmplRealVectorState::new(vec![0.0; from.inner.values.len()]);
91        self.inner
92            .lock()
93            .unwrap()
94            .interpolate(&from.inner, &to.inner, t, &mut result_state);
95        JsRealVectorState {
96            inner: Arc::new(result_state),
97        }
98    }
99
100    #[wasm_bindgen(js_name = getDimension)]
101    pub fn get_dimension(&self) -> usize {
102        self.inner.lock().unwrap().dimension
103    }
104
105    #[wasm_bindgen(js_name = getMaximumExtent)]
106    pub fn get_maximum_extent(&self) -> f64 {
107        self.inner.lock().unwrap().get_maximum_extent()
108    }
109
110    #[wasm_bindgen(js_name = getLongestValidSegmentLength)]
111    pub fn get_longest_valid_segment_length(&self) -> f64 {
112        self.inner
113            .lock()
114            .unwrap()
115            .get_longest_valid_segment_length()
116    }
117
118    #[wasm_bindgen(js_name = setLongestValidLineSegmentFraction)]
119    pub fn set_longest_valid_segment_fraction(&mut self, fraction: f64) {
120        self.inner
121            .lock()
122            .unwrap()
123            .set_longest_valid_segment_fraction(fraction);
124    }
125}