1#![warn(missing_docs)]
21#![forbid(unsafe_code)]
22
23use std::collections::HashMap;
24
25use spectrafit_models::{model_from_str, Model};
26use spectrafit_types::types::{ExprEdge, FitGraphSpec, ModelNodeSpec, ModelTypeStr, ParameterSpec};
27
28pub const SCHEMA_VERSION: &str = "0.1";
33
34#[derive(Debug, Clone, Default)]
43pub struct FitGraphBuilder {
44 nodes: Vec<ModelNodeSpec>,
45 expr_edges: Vec<ExprEdge>,
46}
47
48impl FitGraphBuilder {
49 pub fn new() -> Self {
51 Self {
52 nodes: Vec::new(),
53 expr_edges: Vec::new(),
54 }
55 }
56
57 pub fn available_models() -> Vec<&'static str> {
63 ALL_MODELS.iter().map(|m| m.as_str()).collect()
64 }
65
66 pub fn tie(
71 mut self,
72 target_node: impl Into<String>,
73 target_param: impl Into<String>,
74 expression: impl Into<String>,
75 ) -> Self {
76 self.expr_edges.push(ExprEdge {
77 target_node: target_node.into(),
78 target_param: target_param.into(),
79 expression: expression.into(),
80 });
81 self
82 }
83
84 pub fn build(self) -> FitGraphSpec {
86 FitGraphSpec {
87 schema_version: SCHEMA_VERSION.to_string(),
88 nodes: self.nodes,
89 expr_edges: self.expr_edges,
90 }
91 }
92
93 pub fn add_gaussian(self, id: &str, amplitude: f64, center: f64, sigma: f64) -> Self {
103 self.add_node(id, ModelTypeStr::Gaussian, &[amplitude, center, sigma])
104 }
105
106 pub fn add_gaussian2d(
111 self,
112 id: &str,
113 amplitude: f64,
114 center_x: f64,
115 center_y: f64,
116 sigma_x: f64,
117 sigma_y: f64,
118 ) -> Self {
119 self.add_node(
120 id,
121 ModelTypeStr::Gaussian2D,
122 &[amplitude, center_x, center_y, sigma_x, sigma_y],
123 )
124 }
125
126 pub fn add_gaussian_nd(self, id: &str, amplitude: f64, center_0: f64, sigma_0: f64) -> Self {
134 self.add_node(
135 id,
136 ModelTypeStr::GaussianNd,
137 &[amplitude, center_0, sigma_0],
138 )
139 }
140
141 pub fn add_lorentzian(self, id: &str, amplitude: f64, center: f64, sigma: f64) -> Self {
145 self.add_node(id, ModelTypeStr::Lorentzian, &[amplitude, center, sigma])
146 }
147
148 pub fn add_voigt(
157 self,
158 id: &str,
159 amplitude: f64,
160 center: f64,
161 sigma: f64,
162 fraction: f64,
163 ) -> Self {
164 self.add_node(
165 id,
166 ModelTypeStr::Voigt,
167 &[amplitude, center, sigma, fraction],
168 )
169 }
170
171 pub fn add_constant(self, id: &str, c: f64) -> Self {
173 self.add_node(id, ModelTypeStr::Constant, &[c])
174 }
175
176 pub fn add_linear(self, id: &str, slope: f64, intercept: f64) -> Self {
178 self.add_node(id, ModelTypeStr::Linear, &[slope, intercept])
179 }
180
181 pub fn add_quadratic(self, id: &str, amplitude: f64, center: f64, offset: f64) -> Self {
183 self.add_node(id, ModelTypeStr::Quadratic, &[amplitude, center, offset])
184 }
185
186 pub fn add_arctan_step(self, id: &str, amplitude: f64, center: f64, sigma: f64) -> Self {
190 self.add_node(id, ModelTypeStr::ArctanStep, &[amplitude, center, sigma])
191 }
192
193 pub fn add_tanh_step(self, id: &str, amplitude: f64, center: f64, sigma: f64) -> Self {
197 self.add_node(id, ModelTypeStr::TanhStep, &[amplitude, center, sigma])
198 }
199
200 pub fn add_erfc_step(self, id: &str, amplitude: f64, center: f64, sigma: f64) -> Self {
204 self.add_node(id, ModelTypeStr::ErfcStep, &[amplitude, center, sigma])
205 }
206
207 pub fn add_pseudo_voigt(
211 self,
212 id: &str,
213 amplitude: f64,
214 center: f64,
215 sigma: f64,
216 fraction: f64,
217 ) -> Self {
218 self.add_node(
219 id,
220 ModelTypeStr::PseudoVoigt,
221 &[amplitude, center, sigma, fraction],
222 )
223 }
224
225 pub fn add_fano(self, id: &str, amplitude: f64, center: f64, gamma: f64, q: f64) -> Self {
229 self.add_node(id, ModelTypeStr::Fano, &[amplitude, center, gamma, q])
230 }
231
232 pub fn add_double_exponential(self, id: &str, a1: f64, lam1: f64, a2: f64, lam2: f64) -> Self {
234 self.add_node(id, ModelTypeStr::DoubleExponential, &[a1, lam1, a2, lam2])
235 }
236
237 pub fn add_saturating_exponential(self, id: &str, amplitude: f64, rate: f64) -> Self {
241 self.add_node(id, ModelTypeStr::SaturatingExponential, &[amplitude, rate])
242 }
243
244 pub fn add_power_saturation(self, id: &str, amplitude: f64, rate: f64) -> Self {
248 self.add_node(id, ModelTypeStr::PowerSaturation, &[amplitude, rate])
249 }
250
251 pub fn add_power_law_offset(self, id: &str, amplitude: f64, offset: f64, shape: f64) -> Self {
258 self.add_node(
259 id,
260 ModelTypeStr::PowerLawOffset,
261 &[amplitude, offset, shape],
262 )
263 }
264
265 pub fn add_mgh09_rational(
275 self,
276 id: &str,
277 amplitude: f64,
278 num_lin: f64,
279 den_lin: f64,
280 den_const: f64,
281 ) -> Self {
282 self.add_node(
283 id,
284 ModelTypeStr::Mgh09Rational,
285 &[amplitude, num_lin, den_lin, den_const],
286 )
287 }
288
289 pub fn add_true_voigt(
294 self,
295 id: &str,
296 amplitude: f64,
297 center: f64,
298 sigma: f64,
299 gamma: f64,
300 ) -> Self {
301 self.add_node(
302 id,
303 ModelTypeStr::TrueVoigt,
304 &[amplitude, center, sigma, gamma],
305 )
306 }
307
308 pub fn add_skewed_gaussian(
312 self,
313 id: &str,
314 amplitude: f64,
315 center: f64,
316 sigma: f64,
317 gamma: f64,
318 ) -> Self {
319 self.add_node(
320 id,
321 ModelTypeStr::SkewedGaussian,
322 &[amplitude, center, sigma, gamma],
323 )
324 }
325
326 pub fn add_exp_gaussian(
330 self,
331 id: &str,
332 amplitude: f64,
333 center: f64,
334 sigma: f64,
335 gamma: f64,
336 ) -> Self {
337 self.add_node(
338 id,
339 ModelTypeStr::ExpGaussian,
340 &[amplitude, center, sigma, gamma],
341 )
342 }
343
344 pub fn add_doniach_sunjic(
348 self,
349 id: &str,
350 amplitude: f64,
351 center: f64,
352 sigma: f64,
353 gamma: f64,
354 ) -> Self {
355 self.add_node(
356 id,
357 ModelTypeStr::DoniachSunjic,
358 &[amplitude, center, sigma, gamma],
359 )
360 }
361
362 pub fn add_log_normal(self, id: &str, amplitude: f64, center: f64, sigma: f64) -> Self {
366 self.add_node(id, ModelTypeStr::LogNormal, &[amplitude, center, sigma])
367 }
368
369 pub fn add_pearson7(self, id: &str, amplitude: f64, center: f64, sigma: f64, m: f64) -> Self {
373 self.add_node(id, ModelTypeStr::Pearson7, &[amplitude, center, sigma, m])
374 }
375
376 pub fn add_split_gaussian(
380 self,
381 id: &str,
382 amplitude: f64,
383 center: f64,
384 sigma_l: f64,
385 sigma_r: f64,
386 ) -> Self {
387 self.add_node(
388 id,
389 ModelTypeStr::SplitGaussian,
390 &[amplitude, center, sigma_l, sigma_r],
391 )
392 }
393
394 pub fn add_moffat(self, id: &str, amplitude: f64, center: f64, sigma: f64, beta: f64) -> Self {
398 self.add_node(id, ModelTypeStr::Moffat, &[amplitude, center, sigma, beta])
399 }
400
401 pub fn add_students_t(
405 self,
406 id: &str,
407 amplitude: f64,
408 center: f64,
409 sigma: f64,
410 nu: f64,
411 ) -> Self {
412 self.add_node(id, ModelTypeStr::StudentsT, &[amplitude, center, sigma, nu])
413 }
414
415 #[allow(clippy::too_many_arguments)]
423 pub fn add_split_pearson7(
424 self,
425 id: &str,
426 amplitude: f64,
427 center: f64,
428 sigma_l: f64,
429 sigma_r: f64,
430 m_l: f64,
431 m_r: f64,
432 ) -> Self {
433 self.add_node(
434 id,
435 ModelTypeStr::SplitPearson7,
436 &[amplitude, center, sigma_l, sigma_r, m_l, m_r],
437 )
438 }
439
440 pub fn add_breit_wigner(
444 self,
445 id: &str,
446 amplitude: f64,
447 center: f64,
448 sigma: f64,
449 q: f64,
450 ) -> Self {
451 self.add_node(
452 id,
453 ModelTypeStr::BreitWigner,
454 &[amplitude, center, sigma, q],
455 )
456 }
457
458 pub fn add_asym_ir(self, id: &str, amplitude: f64, center: f64, sigma: f64, k: f64) -> Self {
462 self.add_node(id, ModelTypeStr::AsymIr, &[amplitude, center, sigma, k])
463 }
464
465 pub fn add_harmonic_ir(self, id: &str, amplitude: f64, center: f64, sigma: f64) -> Self {
469 self.add_node(id, ModelTypeStr::HarmonicIr, &[amplitude, center, sigma])
470 }
471
472 pub fn add_tauc(self, id: &str, amplitude: f64, e_gap: f64, exponent: f64) -> Self {
476 self.add_node(id, ModelTypeStr::Tauc, &[amplitude, e_gap, exponent])
477 }
478
479 pub fn add_cauchy_dispersion(self, id: &str, a: f64, b: f64, c: f64) -> Self {
481 self.add_node(id, ModelTypeStr::CauchyDispersion, &[a, b, c])
482 }
483
484 pub fn add_kww(self, id: &str, amplitude: f64, tau: f64, beta: f64) -> Self {
488 self.add_node(id, ModelTypeStr::Kww, &[amplitude, tau, beta])
489 }
490
491 fn add_node(mut self, id: &str, model_type: ModelTypeStr, values: &[f64]) -> Self {
496 let wire = model_type.as_str();
497 let kernel: Box<dyn Model> = model_from_str(wire).unwrap_or_else(|| {
498 panic!(
503 "spectrafit-builder: model_from_str({wire:?}) returned None — \
504 add the kernel registration in spectrafit-models::model_from_str"
505 )
506 });
507 let names = kernel.param_names();
508 debug_assert_eq!(
509 names.len(),
510 values.len(),
511 "builder arity mismatch for {wire}: kernel expects {} params, got {}",
512 names.len(),
513 values.len()
514 );
515 let mut parameters: HashMap<String, ParameterSpec> = HashMap::with_capacity(names.len());
516 for (name, value) in names.iter().zip(values.iter()) {
517 parameters.insert((*name).to_string(), default_parameter(*value));
518 }
519 self.nodes.push(ModelNodeSpec {
520 id: id.to_string(),
521 model_type,
522 parameters,
523 dataset_index: None,
524 });
525 self
526 }
527}
528
529fn default_parameter(value: f64) -> ParameterSpec {
531 ParameterSpec {
532 value,
533 min: f64::NEG_INFINITY,
534 max: f64::INFINITY,
535 vary: true,
536 expr: None,
537 scale: None,
538 }
539}
540
541const ALL_MODELS: &[ModelTypeStr] = &[
546 ModelTypeStr::Gaussian,
547 ModelTypeStr::Gaussian2D,
548 ModelTypeStr::GaussianNd,
549 ModelTypeStr::Lorentzian,
550 ModelTypeStr::Voigt,
551 ModelTypeStr::Constant,
552 ModelTypeStr::Linear,
553 ModelTypeStr::Quadratic,
554 ModelTypeStr::ArctanStep,
555 ModelTypeStr::TanhStep,
556 ModelTypeStr::ErfcStep,
557 ModelTypeStr::PseudoVoigt,
558 ModelTypeStr::Fano,
559 ModelTypeStr::DoubleExponential,
560 ModelTypeStr::SaturatingExponential,
561 ModelTypeStr::TrueVoigt,
562 ModelTypeStr::SkewedGaussian,
563 ModelTypeStr::ExpGaussian,
564 ModelTypeStr::DoniachSunjic,
565 ModelTypeStr::LogNormal,
566 ModelTypeStr::Pearson7,
567 ModelTypeStr::SplitGaussian,
568 ModelTypeStr::Moffat,
569 ModelTypeStr::StudentsT,
570 ModelTypeStr::SplitPearson7,
571 ModelTypeStr::BreitWigner,
572 ModelTypeStr::AsymIr,
573 ModelTypeStr::HarmonicIr,
574 ModelTypeStr::Tauc,
575 ModelTypeStr::CauchyDispersion,
576 ModelTypeStr::Kww,
577 ModelTypeStr::PowerSaturation,
578 ModelTypeStr::PowerLawOffset,
579 ModelTypeStr::Mgh09Rational,
580 ModelTypeStr::RationalCubic,
581 ModelTypeStr::GeneralisedLogistic,
582 ModelTypeStr::ExpOverLinear,
583];
584
585#[cfg(test)]
586mod tests {
587 use super::*;
588
589 #[test]
593 fn available_models_matches_model_from_str() {
594 for key in FitGraphBuilder::available_models() {
595 assert!(
596 model_from_str(key).is_some(),
597 "available_models() reported {key:?} but model_from_str does not know it",
598 );
599 }
600 }
601
602 #[test]
604 fn empty_builder_produces_valid_spec() {
605 let g = FitGraphBuilder::new().build();
606 assert_eq!(g.schema_version, SCHEMA_VERSION);
607 assert!(g.nodes.is_empty());
608 assert!(g.expr_edges.is_empty());
609 }
610
611 #[test]
618 fn every_model_type_str_variant_is_covered_by_all_models() {
619 fn covered(variant: &ModelTypeStr) -> bool {
620 ALL_MODELS.iter().any(|m| m.as_str() == variant.as_str())
621 }
622 let representatives = [
624 ModelTypeStr::Gaussian,
625 ModelTypeStr::Gaussian2D,
626 ModelTypeStr::GaussianNd,
627 ModelTypeStr::Lorentzian,
628 ModelTypeStr::Voigt,
629 ModelTypeStr::Constant,
630 ModelTypeStr::Linear,
631 ModelTypeStr::Quadratic,
632 ModelTypeStr::ArctanStep,
633 ModelTypeStr::TanhStep,
634 ModelTypeStr::ErfcStep,
635 ModelTypeStr::PseudoVoigt,
636 ModelTypeStr::Fano,
637 ModelTypeStr::DoubleExponential,
638 ModelTypeStr::SaturatingExponential,
639 ModelTypeStr::TrueVoigt,
640 ModelTypeStr::SkewedGaussian,
641 ModelTypeStr::ExpGaussian,
642 ModelTypeStr::DoniachSunjic,
643 ModelTypeStr::LogNormal,
644 ModelTypeStr::Pearson7,
645 ModelTypeStr::SplitGaussian,
646 ModelTypeStr::Moffat,
647 ModelTypeStr::StudentsT,
648 ModelTypeStr::SplitPearson7,
649 ModelTypeStr::BreitWigner,
650 ModelTypeStr::AsymIr,
651 ModelTypeStr::HarmonicIr,
652 ModelTypeStr::Tauc,
653 ModelTypeStr::CauchyDispersion,
654 ModelTypeStr::Kww,
655 ModelTypeStr::PowerSaturation,
656 ModelTypeStr::PowerLawOffset,
657 ModelTypeStr::Mgh09Rational,
658 ModelTypeStr::RationalCubic,
659 ModelTypeStr::GeneralisedLogistic,
660 ModelTypeStr::ExpOverLinear,
661 ];
662 for v in &representatives {
666 let _exhaustive: () = match v {
669 ModelTypeStr::Gaussian
670 | ModelTypeStr::Gaussian2D
671 | ModelTypeStr::GaussianNd
672 | ModelTypeStr::Lorentzian
673 | ModelTypeStr::Voigt
674 | ModelTypeStr::Constant
675 | ModelTypeStr::Linear
676 | ModelTypeStr::Quadratic
677 | ModelTypeStr::ArctanStep
678 | ModelTypeStr::TanhStep
679 | ModelTypeStr::ErfcStep
680 | ModelTypeStr::PseudoVoigt
681 | ModelTypeStr::Fano
682 | ModelTypeStr::DoubleExponential
683 | ModelTypeStr::SaturatingExponential
684 | ModelTypeStr::TrueVoigt
685 | ModelTypeStr::SkewedGaussian
686 | ModelTypeStr::ExpGaussian
687 | ModelTypeStr::DoniachSunjic
688 | ModelTypeStr::LogNormal
689 | ModelTypeStr::Pearson7
690 | ModelTypeStr::SplitGaussian
691 | ModelTypeStr::Moffat
692 | ModelTypeStr::StudentsT
693 | ModelTypeStr::SplitPearson7
694 | ModelTypeStr::BreitWigner
695 | ModelTypeStr::AsymIr
696 | ModelTypeStr::HarmonicIr
697 | ModelTypeStr::Tauc
698 | ModelTypeStr::CauchyDispersion
699 | ModelTypeStr::Kww
700 | ModelTypeStr::PowerSaturation
701 | ModelTypeStr::PowerLawOffset
702 | ModelTypeStr::Mgh09Rational
703 | ModelTypeStr::RationalCubic
704 | ModelTypeStr::GeneralisedLogistic
705 | ModelTypeStr::ExpOverLinear => (),
706 };
707 assert!(
708 covered(v),
709 "ModelTypeStr::{v:?} is not present in ALL_MODELS — add it \
710 alongside the corresponding `add_<name>()` fluent method",
711 );
712 }
713 assert_eq!(
714 ALL_MODELS.len(),
715 representatives.len(),
716 "ALL_MODELS length drifted from the exhaustive variant list",
717 );
718 }
719
720 #[test]
722 fn tie_accumulates_edges_in_order() {
723 let g = FitGraphBuilder::new()
724 .tie("g0", "center", "g1.center + 0.5")
725 .tie("g1", "sigma", "g0.sigma")
726 .build();
727 assert_eq!(g.expr_edges.len(), 2);
728 assert_eq!(g.expr_edges[0].target_node, "g0");
729 assert_eq!(g.expr_edges[0].target_param, "center");
730 assert_eq!(g.expr_edges[0].expression, "g1.center + 0.5");
731 assert_eq!(g.expr_edges[1].target_node, "g1");
732 }
733}