Skip to main content

oxmpl_js/geometric/
rrt.rs

1// Copyright (c) 2025 Ross Gardiner, Junior Sundar
2//
3// SPDX-License-Identifier: BSD-3-Clause
4
5use crate::base::{
6    goal::JsGoal,
7    path::JsPath,
8    planner::JsPlannerConfig,
9    problem_definition::{JsProblemDefinition, ProblemDefinitionVariant},
10    state_validity_checker::JsStateValidityChecker,
11};
12use oxmpl::base::{
13    planner::Planner,
14    space::{
15        CompoundStateSpace, RealVectorStateSpace, SE2StateSpace, SE3StateSpace, SO2StateSpace,
16        SO3StateSpace,
17    },
18    state::{CompoundState, RealVectorState, SE2State, SE3State, SO2State, SO3State},
19};
20use oxmpl::geometric::RRT;
21use std::sync::Arc;
22use std::time::Duration;
23use wasm_bindgen::prelude::*;
24
25enum RrtVariant {
26    RealVector(RRT<RealVectorState, RealVectorStateSpace, JsGoal>),
27    SO2(RRT<SO2State, SO2StateSpace, JsGoal>),
28    SO3(RRT<SO3State, SO3StateSpace, JsGoal>),
29    Compound(RRT<CompoundState, CompoundStateSpace, JsGoal>),
30    SE2(RRT<SE2State, SE2StateSpace, JsGoal>),
31    SE3(RRT<SE3State, SE3StateSpace, JsGoal>),
32}
33
34#[wasm_bindgen(js_name = RRT)]
35pub struct JsRRT {
36    planner: RrtVariant,
37    pd: ProblemDefinitionVariant,
38}
39
40#[wasm_bindgen(js_class = RRT)]
41impl JsRRT {
42    #[wasm_bindgen(constructor)]
43    pub fn new(
44        max_distance: f64,
45        goal_bias: f64,
46        problem_def: &JsProblemDefinition,
47        config: &JsPlannerConfig,
48    ) -> Self {
49        let planner_config = config.into();
50        match &problem_def.inner {
51            ProblemDefinitionVariant::RealVector(pd) => Self {
52                planner: RrtVariant::RealVector(RRT::new(max_distance, goal_bias, &planner_config)),
53                pd: ProblemDefinitionVariant::RealVector(pd.clone()),
54            },
55            ProblemDefinitionVariant::SO2(pd) => Self {
56                planner: RrtVariant::SO2(RRT::new(max_distance, goal_bias, &planner_config)),
57                pd: ProblemDefinitionVariant::SO2(pd.clone()),
58            },
59            ProblemDefinitionVariant::SO3(pd) => Self {
60                planner: RrtVariant::SO3(RRT::new(max_distance, goal_bias, &planner_config)),
61                pd: ProblemDefinitionVariant::SO3(pd.clone()),
62            },
63            ProblemDefinitionVariant::Compound(pd) => Self {
64                planner: RrtVariant::Compound(RRT::new(max_distance, goal_bias, &planner_config)),
65                pd: ProblemDefinitionVariant::Compound(pd.clone()),
66            },
67            ProblemDefinitionVariant::SE2(pd) => Self {
68                planner: RrtVariant::SE2(RRT::new(max_distance, goal_bias, &planner_config)),
69                pd: ProblemDefinitionVariant::SE2(pd.clone()),
70            },
71            ProblemDefinitionVariant::SE3(pd) => Self {
72                planner: RrtVariant::SE3(RRT::new(max_distance, goal_bias, &planner_config)),
73                pd: ProblemDefinitionVariant::SE3(pd.clone()),
74            },
75        }
76    }
77
78    pub fn setup(&mut self, validity_checker: &JsStateValidityChecker) {
79        let checker = Arc::new(validity_checker.clone());
80        match &mut self.planner {
81            RrtVariant::RealVector(p) => {
82                if let ProblemDefinitionVariant::RealVector(pd) = &self.pd {
83                    p.setup(pd.clone(), checker);
84                }
85            }
86            RrtVariant::SO2(p) => {
87                if let ProblemDefinitionVariant::SO2(pd) = &self.pd {
88                    p.setup(pd.clone(), checker);
89                }
90            }
91            RrtVariant::SO3(p) => {
92                if let ProblemDefinitionVariant::SO3(pd) = &self.pd {
93                    p.setup(pd.clone(), checker);
94                }
95            }
96            RrtVariant::Compound(p) => {
97                if let ProblemDefinitionVariant::Compound(pd) = &self.pd {
98                    p.setup(pd.clone(), checker);
99                }
100            }
101            RrtVariant::SE2(p) => {
102                if let ProblemDefinitionVariant::SE2(pd) = &self.pd {
103                    p.setup(pd.clone(), checker);
104                }
105            }
106            RrtVariant::SE3(p) => {
107                if let ProblemDefinitionVariant::SE3(pd) = &self.pd {
108                    p.setup(pd.clone(), checker);
109                }
110            }
111        }
112    }
113
114    pub fn solve(&mut self, timeout_secs: f32) -> Result<JsPath, String> {
115        let timeout = Duration::from_secs_f32(timeout_secs);
116        match &mut self.planner {
117            RrtVariant::RealVector(p) => p
118                .solve(timeout)
119                .map(JsPath::from)
120                .map_err(|e| e.to_string()),
121            RrtVariant::SO2(p) => p
122                .solve(timeout)
123                .map(JsPath::from)
124                .map_err(|e| e.to_string()),
125            RrtVariant::SO3(p) => p
126                .solve(timeout)
127                .map(JsPath::from)
128                .map_err(|e| e.to_string()),
129            RrtVariant::Compound(p) => p
130                .solve(timeout)
131                .map(JsPath::from)
132                .map_err(|e| e.to_string()),
133            RrtVariant::SE2(p) => p
134                .solve(timeout)
135                .map(JsPath::from)
136                .map_err(|e| e.to_string()),
137            RrtVariant::SE3(p) => p
138                .solve(timeout)
139                .map(JsPath::from)
140                .map_err(|e| e.to_string()),
141        }
142    }
143}