Skip to main content

sva_engine/
cast.rs

1// Concern: the five written casts, the type each produces and the mismatch it writes | Non-concern: performing the approximation (sva-samples) | IO: (Cast, &[Ty]) -> Ty or Mismatch
2
3use sva_formula::{Codomain, Held, Origin, Ty, Var};
4
5/// The operand types, the subterm that blocked, and the cast to write.
6#[derive(Clone, Debug, PartialEq)]
7pub struct Mismatch {
8    pub code: &'static str,
9    pub got: Vec<Ty>,
10    pub blocker: Option<(Origin, String)>,
11    pub repair: String,
12}
13
14impl Mismatch {
15    pub fn new(code: &'static str, got: &[Ty], repair: impl Into<String>) -> Mismatch {
16        Mismatch {
17            code,
18            got: got.to_vec(),
19            blocker: None,
20            repair: repair.into(),
21        }
22    }
23
24    pub fn blocked_by(mut self, origin: Origin, what: impl Into<String>) -> Mismatch {
25        self.blocker = Some((origin, what.into()));
26        self
27    }
28}
29
30#[derive(Clone, Copy, Debug, PartialEq, Eq)]
31pub enum Cast {
32    Sample,
33    Fourier,
34    IFourier,
35    Stft { window: usize, hop: usize },
36    Istft,
37}
38
39/// The named arguments a call carries, already folded to numbers.
40pub type Named<'a> = [(&'a str, f64)];
41
42impl Cast {
43    pub const NAMES: [&'static str; 5] = ["sample", "fourier", "ifourier", "stft", "istft"];
44
45    pub fn from_name(name: &str, named: &Named) -> Option<Cast> {
46        let read = |key: &str| named.iter().find(|(k, _)| *k == key).map(|(_, v)| *v);
47        Some(match name {
48            "sample" => Cast::Sample,
49            "fourier" => Cast::Fourier,
50            "ifourier" => Cast::IFourier,
51            "stft" => Cast::Stft {
52                window: read("window").unwrap_or(0.0) as usize,
53                hop: read("hop").unwrap_or(0.0) as usize,
54            },
55            "istft" => Cast::Istft,
56            _ => return None,
57        })
58    }
59
60    pub fn name(self) -> &'static str {
61        match self {
62            Cast::Sample => "sample",
63            Cast::Fourier => "fourier",
64            Cast::IFourier => "ifourier",
65            Cast::Stft { .. } => "stft",
66            Cast::Istft => "istft",
67        }
68    }
69
70    /// A window and a hop belong to the cast, so they are checked where the call is written.
71    pub fn check_arguments(self) -> Result<(), Mismatch> {
72        match self {
73            Cast::Stft { window, hop } if window == 0 || hop == 0 => Err(Mismatch::new(
74                "cast.missing_window",
75                &[],
76                "stft needs window= and hop=",
77            )),
78            _ => Ok(()),
79        }
80    }
81
82    pub fn resolve(self, args: &[Ty]) -> Result<Ty, Mismatch> {
83        let [arg] = args else {
84            return Err(Mismatch::new(
85                "grammar.arity",
86                args,
87                format!("`{}` takes one signal", self.name()),
88            ));
89        };
90        match self {
91            Cast::Fourier => self.project(*arg, Var::F, args),
92            Cast::IFourier => self.project(*arg, Var::T, args),
93            Cast::Sample => match arg.held {
94                Held::Form(Var::T) => Ok(sampled(arg)),
95                Held::Form(Var::F) if arg.has_dual() => Ok(sampled(arg)),
96                Held::Form(Var::F) => Err(Mismatch::new(
97                    "cast.left_algebra",
98                    args,
99                    "write sample(ifourier(x)) to sample a spectrum in t",
100                )),
101                Held::Sampled | Held::Frames => Err(Mismatch::new(
102                    "type.samples_in_closed_form",
103                    args,
104                    "this is already samples; drop the sample(...)",
105                )),
106            },
107            Cast::Stft { .. } => match arg.held {
108                Held::Sampled => Ok(Ty {
109                    held: Held::Frames,
110                    ..*arg
111                }),
112                Held::Frames => Err(Mismatch::new(
113                    "cast.stft_needs_samples",
114                    args,
115                    "this is already frames; drop the stft(...)",
116                )),
117                _ => Err(Mismatch::new(
118                    "cast.stft_needs_samples",
119                    args,
120                    "stft reads samples. write stft(sample(x), window=, hop=)",
121                )),
122            },
123            Cast::Istft => match arg.held {
124                Held::Frames => Ok(sampled(arg)),
125                _ => Err(Mismatch::new(
126                    "cast.istft_needs_frames",
127                    args,
128                    "istft reads frames. write istft(stft(x, window=, hop=))",
129                )),
130            },
131        }
132    }
133
134    /// A retyping cast selects the other axis of a value that has both, and changes nothing:
135    /// what it states is the axis the expression above it is written on.
136    fn project(self, arg: Ty, onto: Var, args: &[Ty]) -> Result<Ty, Mismatch> {
137        match arg.held {
138            Held::Form(_) if arg.has_dual() => Ok(Ty {
139                held: Held::Form(onto),
140                codomain: match onto {
141                    Var::F => Codomain::Complex,
142                    Var::T => arg.codomain,
143                },
144                ..arg
145            }),
146            Held::Sampled | Held::Frames => Err(Mismatch::new(
147                "cast.samples_are_terminal",
148                args,
149                "nothing returns from samples to a closed form",
150            )),
151            _ => Err(Mismatch::new(
152                "cast.left_algebra",
153                args,
154                format!(
155                    "write {}(sample(x)) for a short-time spectrum, labeled measured",
156                    self.name()
157                ),
158            )),
159        }
160    }
161}
162
163fn sampled(arg: &Ty) -> Ty {
164    Ty {
165        held: Held::Sampled,
166        dual: false,
167        ..*arg
168    }
169}