use oxmpl::base::{
space::{RealVectorStateSpace as OxmplRealVectorStateSpace, StateSpace},
state::RealVectorState as OxmplRealVectorState,
};
use rand::rng;
use std::sync::{Arc, Mutex};
use wasm_bindgen::prelude::*;
use crate::base::JsRealVectorState;
#[wasm_bindgen(js_name = RealVectorStateSpace)]
pub struct JsRealVectorStateSpace {
#[wasm_bindgen(skip)]
pub inner: Arc<Mutex<OxmplRealVectorStateSpace>>,
}
#[wasm_bindgen(js_class = RealVectorStateSpace)]
impl JsRealVectorStateSpace {
#[wasm_bindgen(constructor)]
pub fn new(
dimension: usize,
bounds: Option<Vec<f64>>,
) -> Result<JsRealVectorStateSpace, String> {
let bounds_vec = if let Some(b) = bounds {
if b.len() != dimension * 2 {
return Err(format!(
"Bounds array must have {} elements (2 per dimension)",
dimension * 2
));
}
let mut bounds_tuples = Vec::new();
for i in 0..dimension {
bounds_tuples.push((b[i * 2], b[i * 2 + 1]));
}
Some(bounds_tuples)
} else {
None
};
match OxmplRealVectorStateSpace::new(dimension, bounds_vec) {
Ok(space) => Ok(Self {
inner: Arc::new(Mutex::new(space)),
}),
Err(e) => Err(e.to_string()),
}
}
#[wasm_bindgen(js_name = sample)]
pub fn sample(&self) -> Result<JsRealVectorState, String> {
let mut rng = rng();
match self.inner.lock().unwrap().sample_uniform(&mut rng) {
Ok(state) => Ok(JsRealVectorState::new(state.values)),
Err(e) => Err(e.to_string()),
}
}
#[wasm_bindgen(js_name = distance)]
pub fn distance(&self, state1: &JsRealVectorState, state2: &JsRealVectorState) -> f64 {
self.inner
.lock()
.unwrap()
.distance(&state1.inner, &state2.inner)
}
#[wasm_bindgen(js_name = satisfiesBounds)]
pub fn satisfies_bounds(&self, state: &JsRealVectorState) -> bool {
self.inner.lock().unwrap().satisfies_bounds(&state.inner)
}
#[wasm_bindgen(js_name = enforceBounds)]
pub fn enforce_bounds(&self, state: &JsRealVectorState) -> JsRealVectorState {
let mut new_state = (*state.inner).clone();
self.inner.lock().unwrap().enforce_bounds(&mut new_state);
JsRealVectorState {
inner: Arc::new(new_state),
}
}
#[wasm_bindgen(js_name = interpolate)]
pub fn interpolate(
&self,
from: &JsRealVectorState,
to: &JsRealVectorState,
t: f64,
) -> JsRealVectorState {
let mut result_state = OxmplRealVectorState::new(vec![0.0; from.inner.values.len()]);
self.inner
.lock()
.unwrap()
.interpolate(&from.inner, &to.inner, t, &mut result_state);
JsRealVectorState {
inner: Arc::new(result_state),
}
}
#[wasm_bindgen(js_name = getDimension)]
pub fn get_dimension(&self) -> usize {
self.inner.lock().unwrap().dimension
}
#[wasm_bindgen(js_name = getMaximumExtent)]
pub fn get_maximum_extent(&self) -> f64 {
self.inner.lock().unwrap().get_maximum_extent()
}
#[wasm_bindgen(js_name = getLongestValidSegmentLength)]
pub fn get_longest_valid_segment_length(&self) -> f64 {
self.inner
.lock()
.unwrap()
.get_longest_valid_segment_length()
}
#[wasm_bindgen(js_name = setLongestValidLineSegmentFraction)]
pub fn set_longest_valid_segment_fraction(&mut self, fraction: f64) {
self.inner
.lock()
.unwrap()
.set_longest_valid_segment_fraction(fraction);
}
}