1use 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}