1use std::{collections::BTreeMap, sync::Arc};
4
5use sim_kernel::{
6 Args, ClassRef, Cx, DefaultFactory, Error, Expr, Factory, HandleStore, Object, Ref, Result,
7 Symbol, Term, Value,
8};
9
10use super::{options::parse_symbolish_value, registry::global_numeric_registry};
11
12#[derive(Clone, Debug, PartialEq, Eq)]
14pub enum PipelineKind {
15 OdeSolve,
17 Quadrature,
19}
20
21impl PipelineKind {
22 pub fn symbol(&self) -> Symbol {
24 match self {
25 Self::OdeSolve => Symbol::new("ode-solve"),
26 Self::Quadrature => Symbol::new("quadrature"),
27 }
28 }
29}
30
31#[derive(Clone, Debug, PartialEq, Eq)]
33pub enum StateKind {
34 F64,
36 Tensor,
38}
39
40impl StateKind {
41 pub fn symbol(&self) -> Symbol {
43 match self {
44 Self::F64 => Symbol::new("f64"),
45 Self::Tensor => Symbol::new("tensor"),
46 }
47 }
48}
49
50#[derive(Clone, Debug)]
52pub struct ComposedPipeline {
53 pub func_ref: Ref,
55 pub kind: PipelineKind,
57 pub method: Symbol,
59 pub state: StateKind,
61}
62
63impl ComposedPipeline {
64 pub fn new(func_ref: Ref, kind: PipelineKind, method: Symbol, state: StateKind) -> Self {
66 Self {
67 func_ref,
68 kind,
69 method,
70 state,
71 }
72 }
73
74 pub fn table_value(&self, factory: &dyn Factory) -> Result<Value> {
76 factory.table(vec![
77 (
78 Symbol::new("kind"),
79 factory.string("composed-pipeline".to_owned())?,
80 ),
81 (Symbol::new("domain"), factory.symbol(self.kind.symbol())?),
82 (Symbol::new("method"), factory.symbol(self.method.clone())?),
83 (Symbol::new("state"), factory.symbol(self.state.symbol())?),
84 (
85 Symbol::new("func"),
86 factory.expr(Term::Ref(self.func_ref.clone()).into())?,
87 ),
88 ])
89 }
90}
91
92impl Object for ComposedPipeline {
93 fn display(&self, _cx: &mut Cx) -> Result<String> {
94 Ok(format!(
95 "#<composed-pipeline {} {} {}>",
96 self.kind.symbol(),
97 self.method,
98 self.state.symbol()
99 ))
100 }
101
102 fn as_any(&self) -> &dyn std::any::Any {
103 self
104 }
105}
106
107impl sim_kernel::ObjectCompat for ComposedPipeline {
108 fn class(&self, cx: &mut Cx) -> Result<ClassRef> {
109 if let Some(value) = cx
110 .registry()
111 .class_by_symbol(&Symbol::qualified("core", "Table"))
112 {
113 return Ok(value.clone());
114 }
115 DefaultFactory.class_stub(
116 sim_kernel::CORE_TABLE_CLASS_ID,
117 Symbol::qualified("core", "Table"),
118 )
119 }
120
121 fn as_expr(&self, cx: &mut Cx) -> Result<Expr> {
122 self.as_table(cx)?.object().as_expr(cx)
123 }
124
125 fn as_table(&self, cx: &mut Cx) -> Result<Value> {
126 self.table_value(cx.factory())
127 }
128}
129
130pub fn call_numeric_compose(cx: &mut Cx, args: Args) -> Result<Value> {
131 let values = args.into_vec();
132 let pipeline = compose_from_values(cx, &values)?;
133 pipeline_value(cx, pipeline)
134}
135
136pub fn call_numeric_compose_exprs(cx: &mut Cx, args: Vec<Expr>) -> Result<Value> {
137 let pipeline = compose_from_exprs(cx, &args)?;
138 pipeline_value(cx, pipeline)
139}
140
141pub fn call_numeric_run_composed(cx: &mut Cx, args: Args) -> Result<Value> {
142 super::pipeline_run::call_numeric_run_composed(cx, args)
143}
144
145pub fn call_numeric_run_composed_exprs(cx: &mut Cx, args: Vec<Expr>) -> Result<Value> {
146 super::pipeline_run::call_numeric_run_composed_exprs(cx, args)
147}
148
149fn compose_from_values(cx: &mut Cx, values: &[Value]) -> Result<ComposedPipeline> {
150 match values {
151 [func, kind, method, state] if !is_compose_key_value(cx, kind)? => {
152 let func_ref = require_callable_ref(cx, "numeric/compose", func)?;
153 let kind = require_pipeline_kind_value(cx, "numeric/compose", kind)?;
154 let method = require_symbol_value(cx, "numeric/compose", method)?;
155 let state = require_state_kind_value(cx, "numeric/compose", state)?;
156 finish_compose(func_ref, kind, method, state)
157 }
158 [func, rest @ ..] if rest.len().is_multiple_of(2) => {
159 let func_ref = require_callable_ref(cx, "numeric/compose", func)?;
160 let mut options = BTreeMap::<String, Value>::new();
161 for pair in rest.chunks(2) {
162 let key = require_compose_key_value(cx, &pair[0])?;
163 options.insert(key, pair[1].clone());
164 }
165 let kind = require_compose_kind_value(cx, &options)?;
166 let method = require_compose_symbol_value(cx, &options, "method")?;
167 let state = require_compose_state_value(cx, &options)?;
168 finish_compose(func_ref, kind, method, state)
169 }
170 _ => Err(Error::Eval(
171 "numeric/compose expects func, kind, method, state or keyword pairs".to_owned(),
172 )),
173 }
174}
175
176fn compose_from_exprs(cx: &mut Cx, args: &[Expr]) -> Result<ComposedPipeline> {
177 let Some((func_expr, rest)) = args.split_first() else {
178 return Err(Error::Eval(
179 "numeric/compose expects func, kind, method, state or keyword pairs".to_owned(),
180 ));
181 };
182 let func = cx.eval_expr(func_expr.clone())?;
183 let func_ref = require_callable_ref(cx, "numeric/compose", &func)?;
184 if let [kind_expr, method_expr, state_expr] = rest
185 && !is_compose_key_expr(kind_expr)
186 {
187 let kind = require_pipeline_kind_expr("numeric/compose", kind_expr)?;
188 let method = require_symbol_expr("numeric/compose", method_expr)?;
189 let state = require_state_kind_expr("numeric/compose", state_expr)?;
190 return finish_compose(func_ref, kind, method, state);
191 }
192 if !rest.len().is_multiple_of(2) {
193 return Err(Error::Eval(
194 "numeric/compose keyword arguments must be key/value pairs".to_owned(),
195 ));
196 }
197 let mut options = BTreeMap::<String, Symbol>::new();
198 for pair in rest.chunks(2) {
199 options.insert(
200 require_compose_key_expr(&pair[0])?,
201 require_symbol_expr("numeric/compose", &pair[1])?,
202 );
203 }
204 let kind = require_compose_kind_symbol(&options)?;
205 let method = require_compose_symbol(&options, "method")?;
206 let state = parse_state_kind(&require_compose_symbol(&options, "state")?).ok_or_else(|| {
207 Error::Eval("numeric/compose expected state kind f64 or tensor".to_owned())
208 })?;
209 finish_compose(func_ref, kind, method, state)
210}
211
212fn require_callable_ref(cx: &mut Cx, name: &str, value: &Value) -> Result<Ref> {
213 value.object().as_callable().ok_or_else(|| {
214 Error::Eval(format!(
215 "{name} expects its first argument to be a Func or ordinary callable value"
216 ))
217 })?;
218 Ok(Ref::Handle(cx.handles_mut().intern(value.clone())))
219}
220
221fn pipeline_value(cx: &mut Cx, pipeline: ComposedPipeline) -> Result<Value> {
222 cx.factory().opaque(Arc::new(pipeline))
223}
224
225fn require_pipeline_kind_value(cx: &mut Cx, name: &str, value: &Value) -> Result<PipelineKind> {
226 let symbol = require_symbol_value(cx, name, value)?;
227 parse_pipeline_kind(&symbol).ok_or_else(|| {
228 Error::Eval(format!(
229 "{name} expected pipeline kind ode-solve or quadrature"
230 ))
231 })
232}
233
234fn require_state_kind_value(cx: &mut Cx, name: &str, value: &Value) -> Result<StateKind> {
235 let symbol = require_symbol_value(cx, name, value)?;
236 parse_state_kind(&symbol)
237 .ok_or_else(|| Error::Eval(format!("{name} expected state kind f64 or tensor")))
238}
239
240fn require_symbol_value(cx: &mut Cx, name: &str, value: &Value) -> Result<Symbol> {
241 parse_symbolish_value(cx, value)?
242 .ok_or_else(|| Error::Eval(format!("{name} expected a symbol argument")))
243}
244
245fn finish_compose(
246 func_ref: Ref,
247 kind: PipelineKind,
248 method: Symbol,
249 state: StateKind,
250) -> Result<ComposedPipeline> {
251 if kind == PipelineKind::Quadrature {
252 validate_quadrature_method(&method)?;
253 }
254 Ok(ComposedPipeline::new(func_ref, kind, method, state))
255}
256
257fn require_compose_kind_value(
258 cx: &mut Cx,
259 options: &BTreeMap<String, Value>,
260) -> Result<PipelineKind> {
261 let symbol = options
262 .get("domain")
263 .or_else(|| options.get("kind"))
264 .ok_or_else(|| Error::Eval("numeric/compose missing :domain".to_owned()))
265 .and_then(|value| require_symbol_value(cx, "numeric/compose", value))?;
266 parse_pipeline_kind(&symbol).ok_or_else(|| {
267 Error::Eval("numeric/compose expected domain ode-solve or quadrature".to_owned())
268 })
269}
270
271fn require_compose_kind_symbol(options: &BTreeMap<String, Symbol>) -> Result<PipelineKind> {
272 let symbol = options
273 .get("domain")
274 .or_else(|| options.get("kind"))
275 .ok_or_else(|| Error::Eval("numeric/compose missing :domain".to_owned()))?;
276 parse_pipeline_kind(symbol).ok_or_else(|| {
277 Error::Eval("numeric/compose expected domain ode-solve or quadrature".to_owned())
278 })
279}
280
281fn require_compose_state_value(
282 cx: &mut Cx,
283 options: &BTreeMap<String, Value>,
284) -> Result<StateKind> {
285 let symbol = require_compose_symbol_value(cx, options, "state")?;
286 parse_state_kind(&symbol)
287 .ok_or_else(|| Error::Eval("numeric/compose expected state kind f64 or tensor".to_owned()))
288}
289
290fn require_compose_symbol_value(
291 cx: &mut Cx,
292 options: &BTreeMap<String, Value>,
293 key: &str,
294) -> Result<Symbol> {
295 let value = options
296 .get(key)
297 .ok_or_else(|| Error::Eval(format!("numeric/compose missing :{key}")))?;
298 require_symbol_value(cx, "numeric/compose", value)
299}
300
301fn require_compose_symbol(options: &BTreeMap<String, Symbol>, key: &str) -> Result<Symbol> {
302 options
303 .get(key)
304 .cloned()
305 .ok_or_else(|| Error::Eval(format!("numeric/compose missing :{key}")))
306}
307
308fn require_pipeline_kind_expr(name: &str, expr: &Expr) -> Result<PipelineKind> {
309 let symbol = require_symbol_expr(name, expr)?;
310 parse_pipeline_kind(&symbol).ok_or_else(|| {
311 Error::Eval(format!(
312 "{name} expected pipeline kind ode-solve or quadrature"
313 ))
314 })
315}
316
317fn require_state_kind_expr(name: &str, expr: &Expr) -> Result<StateKind> {
318 let symbol = require_symbol_expr(name, expr)?;
319 parse_state_kind(&symbol)
320 .ok_or_else(|| Error::Eval(format!("{name} expected state kind f64 or tensor")))
321}
322
323fn require_symbol_expr(name: &str, expr: &Expr) -> Result<Symbol> {
324 match expr {
325 Expr::Symbol(symbol) => Ok(symbol.clone()),
326 Expr::Quote { expr, .. } => match expr.as_ref() {
327 Expr::Symbol(symbol) => Ok(symbol.clone()),
328 _ => Err(Error::Eval(format!("{name} expected a symbol argument"))),
329 },
330 _ => Err(Error::Eval(format!("{name} expected a symbol argument"))),
331 }
332}
333
334fn parse_pipeline_kind(symbol: &Symbol) -> Option<PipelineKind> {
335 match keyword_name(symbol).as_str() {
336 "ode-solve" => Some(PipelineKind::OdeSolve),
337 "quadrature" => Some(PipelineKind::Quadrature),
338 _ => None,
339 }
340}
341
342fn parse_state_kind(symbol: &Symbol) -> Option<StateKind> {
343 match keyword_name(symbol).as_str() {
344 "f64" => Some(StateKind::F64),
345 "tensor" => Some(StateKind::Tensor),
346 _ => None,
347 }
348}
349
350fn keyword_name(symbol: &Symbol) -> String {
351 symbol
352 .name
353 .strip_prefix(':')
354 .unwrap_or(&symbol.name)
355 .to_owned()
356}
357
358fn is_compose_key_value(cx: &mut Cx, value: &Value) -> Result<bool> {
359 Ok(parse_symbolish_value(cx, value)?
360 .as_ref()
361 .is_some_and(|symbol| is_compose_key_name(&keyword_name(symbol))))
362}
363
364fn require_compose_key_value(cx: &mut Cx, value: &Value) -> Result<String> {
365 parse_symbolish_value(cx, value)?
366 .map(|symbol| keyword_name(&symbol))
367 .filter(|key| is_compose_key_name(key))
368 .ok_or_else(|| Error::Eval("numeric/compose expected keyword argument".to_owned()))
369}
370
371fn is_compose_key_expr(expr: &Expr) -> bool {
372 let Expr::Symbol(symbol) = expr else {
373 return false;
374 };
375 is_compose_key_name(&keyword_name(symbol))
376}
377
378fn require_compose_key_expr(expr: &Expr) -> Result<String> {
379 let Expr::Symbol(symbol) = expr else {
380 return Err(Error::Eval(
381 "numeric/compose expected keyword argument".to_owned(),
382 ));
383 };
384 let key = keyword_name(symbol);
385 if is_compose_key_name(&key) {
386 Ok(key)
387 } else {
388 Err(Error::Eval(format!(
389 "numeric/compose: unknown option :{key}"
390 )))
391 }
392}
393
394fn is_compose_key_name(key: &str) -> bool {
395 matches!(key, "domain" | "kind" | "method" | "state")
396}
397
398fn validate_quadrature_method(method: &Symbol) -> Result<()> {
399 let method = resolve_quad_method(method);
400 let registry = global_numeric_registry()
401 .read()
402 .map_err(|_| Error::PoisonedLock("numeric registry"))?;
403 if registry.quadrature_fixed(&method).is_some()
404 || registry.quadrature_adaptive(&method).is_some()
405 {
406 Ok(())
407 } else {
408 Err(unknown_numeric_method("quadrature", &method))
409 }
410}
411
412fn resolve_quad_method(method: &Symbol) -> Symbol {
413 if *method != Symbol::new("auto") {
414 return method.clone();
415 }
416 Symbol::new("simpson")
417}
418
419fn unknown_numeric_method(kind: &str, method: &Symbol) -> Error {
420 Error::Eval(format!("UnknownNumericMethod: {kind} method {method}"))
421}