1use sva_formula::{Codomain, Held, MAX_WIDTH, Mismatch as MeetMismatch, Ty, Var};
4
5use crate::cast::{Cast, Mismatch};
6use crate::error::{Diagnostic, Located};
7use crate::vocabulary::recognized_named;
8
9#[derive(Clone, Copy, Debug, PartialEq, Eq)]
10pub enum ParamKind {
11 Signal,
12 Scalar,
13 Index,
14 Bound,
15}
16
17#[derive(Clone, Copy, Debug, PartialEq, Eq)]
18pub struct Param {
19 pub name: &'static str,
20 pub required: bool,
21 pub kind: ParamKind,
22}
23
24pub struct Signature {
25 pub name: &'static str,
26 pub params: &'static [Param],
27 pub variadic: bool,
29 pub result: fn(&[Ty]) -> Result<Ty, Mismatch>,
30}
31
32impl Signature {
33 pub fn named(&self) -> &'static [&'static str] {
35 recognized_named(self.name).unwrap_or(&[])
36 }
37
38 pub fn required(&self) -> usize {
39 self.params.iter().filter(|p| p.required).count()
40 }
41}
42
43const fn need(name: &'static str, kind: ParamKind) -> Param {
44 Param {
45 name,
46 required: true,
47 kind,
48 }
49}
50
51const fn opt(name: &'static str, kind: ParamKind) -> Param {
52 Param {
53 name,
54 required: false,
55 kind,
56 }
57}
58
59const SIGNAL: &[Param] = &[need("x", ParamKind::Signal)];
60const TWO: &[Param] = &[need("a", ParamKind::Signal), need("b", ParamKind::Signal)];
61const WAVE: &[Param] = &[
62 need("hz", ParamKind::Scalar),
63 opt("phase", ParamKind::Scalar),
64];
65const CROP: &[Param] = &[
66 need("x", ParamKind::Signal),
67 need("start", ParamKind::Scalar),
68 need("end", ParamKind::Scalar),
69];
70const DRIVEN: &[Param] = &[
71 need("x", ParamKind::Signal),
72 opt("drive", ParamKind::Scalar),
73];
74const RAND: &[Param] = &[
75 opt("key", ParamKind::Signal),
76 opt("seed", ParamKind::Scalar),
77];
78const SERIES_PARAMS: &[Param] = &[
79 need("k", ParamKind::Index),
80 need("lo", ParamKind::Bound),
81 need("hi", ParamKind::Bound),
82 need("term", ParamKind::Signal),
83];
84const CHANNEL_PARAMS: &[Param] = &[need("x", ParamKind::Signal), need("k", ParamKind::Index)];
85const FUNDAMENTAL: &[Param] = &[need("f0", ParamKind::Scalar)];
86const SEED: &[Param] = &[need("seed", ParamKind::Scalar)];
87const LENGTH: &[Param] = &[need("length", ParamKind::Scalar)];
88const VELOCITY: &[Param] = &[need("vel", ParamKind::Scalar)];
89const VOLUME: &[Param] = &[need("volume", ParamKind::Scalar)];
90const RECTANGLE: &[Param] = &[need("lx", ParamKind::Scalar), opt("ly", ParamKind::Scalar)];
91const BOX_SIDES: &[Param] = &[
92 need("lx", ParamKind::Scalar),
93 opt("ly", ParamKind::Scalar),
94 opt("lz", ParamKind::Scalar),
95];
96const FILTER_ONE_POLE: &[Param] = &[
97 need("x", ParamKind::Signal),
98 need("cutoff", ParamKind::Scalar),
99];
100const FILTER_Q: &[Param] = &[
101 need("x", ParamKind::Signal),
102 need("cutoff", ParamKind::Scalar),
103 opt("q", ParamKind::Scalar),
104];
105const FILTER_GAIN: &[Param] = &[
106 need("x", ParamKind::Signal),
107 need("cutoff", ParamKind::Scalar),
108 opt("q", ParamKind::Scalar),
109 opt("gain", ParamKind::Scalar),
110];
111
112fn neutral() -> Ty {
115 Ty::form(Var::T, true, Codomain::Real)
116}
117
118fn fault(m: MeetMismatch, args: &[Ty]) -> Mismatch {
119 let (code, repair) = match m {
120 MeetMismatch::Domain => (
121 "type.domain_mismatch",
122 "write ifourier on the f side to work in t, or fourier on the t side to work in f",
123 ),
124 MeetMismatch::SamplesInClosedForm => (
125 "type.samples_in_closed_form",
126 "write sample(...) on the closed form to move the whole expression into samples",
127 ),
128 MeetMismatch::Rate => ("type.rate_conflict", "one rate per expression"),
129 MeetMismatch::Frames => ("type.frame_mismatch", "one window and hop per expression"),
130 MeetMismatch::Width => (
131 "type.width_mismatch",
132 "give both sides the same width, or make one of them mono",
133 ),
134 };
135 Mismatch::new(code, args, repair)
136}
137
138fn meet(args: &[Ty]) -> Result<Ty, Mismatch> {
140 let Some((first, rest)) = args.split_first() else {
141 return Ok(neutral());
142 };
143 let mut acc = *first;
144 for ty in rest {
145 acc = acc.meet(*ty).map_err(|m| fault(m, args))?;
146 }
147 Ok(acc)
148}
149
150fn scalar(_args: &[Ty]) -> Result<Ty, Mismatch> {
151 Ok(neutral())
152}
153
154fn dual_form(_args: &[Ty]) -> Result<Ty, Mismatch> {
155 Ok(neutral())
156}
157
158fn samples(_args: &[Ty]) -> Result<Ty, Mismatch> {
159 Ok(Ty::discrete(Held::Sampled, Codomain::Real))
160}
161
162fn elementwise(args: &[Ty]) -> Result<Ty, Mismatch> {
165 meet(args)
166}
167
168fn filter(args: &[Ty]) -> Result<Ty, Mismatch> {
169 let signal = args.first().copied().unwrap_or_else(neutral);
170 match signal.has_dual() || signal.held == Held::Sampled {
171 true => meet(std::slice::from_ref(&signal)),
172 false => Err(Mismatch::new(
173 "type.filter_needs_a_dual",
174 args,
175 "write the filter over sample(x) to filter at the render rate",
176 )),
177 }
178}
179
180fn series(args: &[Ty]) -> Result<Ty, Mismatch> {
181 Ok(args.get(3).copied().unwrap_or_else(neutral))
182}
183
184fn joined(args: &[Ty]) -> Result<Ty, Mismatch> {
185 let mut width = 0u32;
186 let mono: Vec<Ty> = args.iter().map(|t| Ty { width: 1, ..*t }).collect();
187 for ty in args {
188 width += u32::from(ty.width);
189 }
190 let acc = meet(&mono)?;
191 let Ok(width) = u8::try_from(width) else {
192 return Err(too_wide(args, width));
193 };
194 if width > MAX_WIDTH {
195 return Err(too_wide(args, u32::from(width)));
196 }
197 Ok(Ty { width, ..acc })
198}
199
200fn too_wide(args: &[Ty], width: u32) -> Mismatch {
201 Mismatch::new(
202 "type.width_mismatch",
203 args,
204 format!("{width} components join; a value carries at most {MAX_WIDTH}"),
205 )
206}
207
208fn channel(args: &[Ty]) -> Result<Ty, Mismatch> {
209 let x = args.first().copied().unwrap_or_else(neutral);
210 Ok(Ty { width: 1, ..x })
211}
212
213fn cast_of(name: &str, args: &[Ty]) -> Result<Ty, Mismatch> {
216 Cast::from_name(name, &[("window", 1.0), ("hop", 1.0)])
217 .expect("a cast signature names a cast")
218 .resolve(args)
219}
220
221const fn plain(
222 name: &'static str,
223 params: &'static [Param],
224 result: fn(&[Ty]) -> Result<Ty, Mismatch>,
225) -> Signature {
226 Signature {
227 name,
228 params,
229 variadic: false,
230 result,
231 }
232}
233
234pub const OPERATORS: [&str; 5] = ["+", "-", "*", "/", "%"];
235
236pub const FINITE_DIFFERENCE: [&str; 6] = [
238 "chaigne_askenfelt",
239 "willemsen_bilbao_serafin",
240 "darabundit_scavone",
241 "rhaouti_chaigne_joly",
242 "chaigne_doutaut",
243 "botteldooren",
244];
245
246pub static SIGNATURES: &[Signature] = &[
247 plain("+", TWO, elementwise),
248 plain("-", TWO, elementwise),
249 plain("*", TWO, elementwise),
250 plain("/", TWO, elementwise),
251 plain("%", TWO, elementwise),
252 plain("sin", SIGNAL, elementwise),
253 plain("cos", SIGNAL, elementwise),
254 plain("exp", SIGNAL, elementwise),
255 plain("log", SIGNAL, elementwise),
256 plain("sqrt", SIGNAL, elementwise),
257 plain("abs", SIGNAL, elementwise),
258 plain("tanh", SIGNAL, elementwise),
259 plain("sat", DRIVEN, elementwise),
260 plain("pow", TWO, elementwise),
261 plain("max", TWO, elementwise),
262 plain("min", TWO, elementwise),
263 plain("saw", WAVE, dual_form),
264 plain("square", WAVE, dual_form),
265 plain("triangle", WAVE, dual_form),
266 plain("crop", CROP, elementwise),
267 plain("noise", SEED, dual_form),
268 plain("string", FUNDAMENTAL, dual_form),
269 plain("membrane", RECTANGLE, dual_form),
270 plain("bar", LENGTH, dual_form),
271 plain("bore", LENGTH, dual_form),
272 plain("room", BOX_SIDES, dual_form),
273 plain("hammer_pulse", VELOCITY, dual_form),
274 plain("helmholtz", VOLUME, dual_form),
275 plain("rand", RAND, scalar),
276 plain("delta", SIGNAL, dual_form),
277 plain("pv", SIGNAL, dual_form),
278 plain("sum", SERIES_PARAMS, series),
279 Signature {
280 name: "join",
281 params: TWO,
282 variadic: true,
283 result: joined,
284 },
285 plain("ch", CHANNEL_PARAMS, channel),
286 plain("chaigne_askenfelt", FUNDAMENTAL, samples),
287 plain("willemsen_bilbao_serafin", FUNDAMENTAL, samples),
288 plain("darabundit_scavone", LENGTH, samples),
289 plain("rhaouti_chaigne_joly", FUNDAMENTAL, samples),
290 plain("chaigne_doutaut", FUNDAMENTAL, samples),
291 plain("botteldooren", FUNDAMENTAL, samples),
292 plain("lp", FILTER_ONE_POLE, filter),
293 plain("lowpass", FILTER_Q, filter),
294 plain("highpass", FILTER_Q, filter),
295 plain("bandpass", FILTER_Q, filter),
296 plain("notch", FILTER_Q, filter),
297 plain("peaking", FILTER_GAIN, filter),
298 plain("lowshelf", FILTER_GAIN, filter),
299 plain("highshelf", FILTER_GAIN, filter),
300 plain("sample", SIGNAL, |a| cast_of("sample", a)),
301 plain("fourier", SIGNAL, |a| cast_of("fourier", a)),
302 plain("ifourier", SIGNAL, |a| cast_of("ifourier", a)),
303 plain("stft", SIGNAL, |a| cast_of("stft", a)),
304 plain("istft", SIGNAL, |a| cast_of("istft", a)),
305];
306
307pub fn signature(name: &str) -> Option<&'static Signature> {
308 SIGNATURES.iter().find(|s| s.name == name)
309}
310
311pub fn resolve(name: &str, args: &[Ty]) -> Result<Ty, Mismatch> {
314 let Some(sig) = signature(name) else {
315 return Err(Mismatch::new(
316 "grammar.unknown_name",
317 args,
318 format!("`{name}` is not a builtin, a unit or a note name"),
319 ));
320 };
321 (sig.result)(args)
322}
323
324pub fn check_arity(name: &str, positional: usize, named: &[String]) -> Result<(), Mismatch> {
327 let Some(sig) = signature(name) else {
328 return Ok(());
329 };
330 let ceiling = if sig.variadic {
331 usize::from(MAX_WIDTH)
332 } else {
333 sig.params.len()
334 };
335 let mut by_name = 0;
336 for key in named {
337 let Some(at) = sig.params.iter().position(|p| p.name == *key) else {
338 if sig.named().contains(&key.as_str()) {
339 continue;
340 }
341 return Err(Mismatch::new(
342 "grammar.unknown_named_argument",
343 &[],
344 format!("`{name}` does not read `{key}`"),
345 ));
346 };
347 let param = sig.params[at];
348 if param.kind == ParamKind::Signal {
349 return Err(Mismatch::new(
350 "grammar.arity",
351 &[],
352 format!("give the signal `{name}` reads by position, not as `{key}=`"),
353 ));
354 }
355 if at < positional {
356 return Err(Mismatch::new(
357 "grammar.arity",
358 &[],
359 format!("give `{key}` to `{name}` by position or by name, not both"),
360 ));
361 }
362 by_name += usize::from(param.required);
363 }
364 match positional + by_name >= sig.required() && positional <= ceiling {
365 true => Ok(()),
366 false => Err(Mismatch::new(
367 "grammar.arity",
368 &[],
369 format!("`{name}` takes {} arguments", sig.params.len()),
370 )),
371 }
372}
373
374pub fn notation(ty: Ty) -> &'static str {
376 match (ty.held, ty.has_dual()) {
377 (Held::Form(Var::T), false) => "Form(t)",
378 (Held::Form(Var::T), true) => "Form(t) dual",
379 (Held::Form(Var::F), false) => "Form(f)",
380 (Held::Form(Var::F), true) => "Form(f) dual",
381 (Held::Sampled, _) => "Sampled",
382 (Held::Frames, _) => "Frames",
383 }
384}
385
386pub fn describe(ty: Ty) -> String {
388 let word = match ty.held {
389 Held::Form(var) => format!(
390 "a closed form in `{}` with {} dual",
391 var.as_str(),
392 match ty.has_dual() {
393 true => "a",
394 false => "no",
395 }
396 ),
397 Held::Sampled => "samples".to_string(),
398 Held::Frames => "frames".to_string(),
399 };
400 match (ty.held, ty.rate, ty.width) {
401 (Held::Sampled, Some(rate), _) => format!("samples at {rate} Hz"),
402 (_, _, 1) => word,
403 (_, _, width) => format!("{word}, {width} components"),
404 }
405}
406
407pub fn format_refusal(
410 call: &str,
411 m: &Mismatch,
412 at: Located,
413 operands: &[(String, Ty)],
414) -> Diagnostic {
415 let types: Vec<String> = m.got.iter().map(|t| describe(*t)).collect();
416 let named: Vec<String> = operands
417 .iter()
418 .map(|(name, ty)| format!("{name}: {}", describe(*ty)))
419 .collect();
420 let head = match &m.blocker {
421 Some((_, what)) => {
422 format!("`{call}` refused: subterm {what} left A")
423 }
424 None => format!("`{call}` has no overload for ({}).", types.join(", ")),
425 };
426 let message = match named.is_empty() {
427 true => head,
428 false => format!("{head} {}.", named.join(". ")),
429 };
430 Diagnostic {
431 code: m.code.to_string(),
432 message,
433 location: at,
434 help: m.repair.clone(),
435 }
436}