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 need("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)
217 .map(|cast| match cast {
218 Cast::Stft { .. } => Cast::Stft { window: 1, hop: 1 },
219 other => other,
220 })
221 .expect("a cast signature names a cast")
222 .resolve(args)
223}
224
225const fn plain(
226 name: &'static str,
227 params: &'static [Param],
228 result: fn(&[Ty]) -> Result<Ty, Mismatch>,
229) -> Signature {
230 Signature {
231 name,
232 params,
233 variadic: false,
234 result,
235 }
236}
237
238pub const OPERATORS: [&str; 5] = ["+", "-", "*", "/", "%"];
239
240pub const FINITE_DIFFERENCE: [&str; 6] = [
242 "chaigne_askenfelt",
243 "willemsen_bilbao_serafin",
244 "darabundit_scavone",
245 "rhaouti_chaigne_joly",
246 "chaigne_doutaut",
247 "botteldooren",
248];
249
250pub static SIGNATURES: &[Signature] = &[
251 plain("+", TWO, elementwise),
252 plain("-", TWO, elementwise),
253 plain("*", TWO, elementwise),
254 plain("/", TWO, elementwise),
255 plain("%", TWO, elementwise),
256 plain("sin", SIGNAL, elementwise),
257 plain("cos", SIGNAL, elementwise),
258 plain("exp", SIGNAL, elementwise),
259 plain("log", SIGNAL, elementwise),
260 plain("sqrt", SIGNAL, elementwise),
261 plain("abs", SIGNAL, elementwise),
262 plain("tanh", SIGNAL, elementwise),
263 plain("step", SIGNAL, elementwise),
264 plain("sat", DRIVEN, elementwise),
265 plain("pow", TWO, elementwise),
266 plain("max", TWO, elementwise),
267 plain("min", TWO, elementwise),
268 plain("saw", WAVE, dual_form),
269 plain("square", WAVE, dual_form),
270 plain("triangle", WAVE, dual_form),
271 plain("crop", CROP, elementwise),
272 plain("noise", SEED, dual_form),
273 plain("string", FUNDAMENTAL, dual_form),
274 plain("membrane", RECTANGLE, dual_form),
275 plain("bar", LENGTH, dual_form),
276 plain("bore", LENGTH, dual_form),
277 plain("room", BOX_SIDES, dual_form),
278 plain("hammer_pulse", VELOCITY, dual_form),
279 plain("helmholtz", VOLUME, dual_form),
280 plain("rand", RAND, scalar),
281 plain("delta", SIGNAL, dual_form),
282 plain("pv", SIGNAL, dual_form),
283 plain("sum", SERIES_PARAMS, series),
284 Signature {
285 name: "join",
286 params: TWO,
287 variadic: true,
288 result: joined,
289 },
290 plain("ch", CHANNEL_PARAMS, channel),
291 plain("chaigne_askenfelt", FUNDAMENTAL, samples),
292 plain("willemsen_bilbao_serafin", FUNDAMENTAL, samples),
293 plain("darabundit_scavone", LENGTH, samples),
294 plain("rhaouti_chaigne_joly", FUNDAMENTAL, samples),
295 plain("chaigne_doutaut", FUNDAMENTAL, samples),
296 plain("botteldooren", FUNDAMENTAL, samples),
297 plain("lp", FILTER_ONE_POLE, filter),
298 plain("lowpass", FILTER_Q, filter),
299 plain("highpass", FILTER_Q, filter),
300 plain("bandpass", FILTER_Q, filter),
301 plain("notch", FILTER_Q, filter),
302 plain("peaking", FILTER_GAIN, filter),
303 plain("lowshelf", FILTER_GAIN, filter),
304 plain("highshelf", FILTER_GAIN, filter),
305 plain("sample", SIGNAL, |a| cast_of("sample", a)),
306 plain("fourier", SIGNAL, |a| cast_of("fourier", a)),
307 plain("ifourier", SIGNAL, |a| cast_of("ifourier", a)),
308 plain("stft", SIGNAL, |a| cast_of("stft", a)),
309 plain("istft", SIGNAL, |a| cast_of("istft", a)),
310];
311
312pub fn signature(name: &str) -> Option<&'static Signature> {
313 SIGNATURES.iter().find(|s| s.name == name)
314}
315
316pub fn resolve(name: &str, args: &[Ty]) -> Result<Ty, Mismatch> {
320 let Some(sig) = signature(name) else {
321 return Err(Mismatch::new(
322 "grammar.unknown_name",
323 args,
324 format!("`{name}` is not a builtin, a unit or a note name"),
325 ));
326 };
327 (sig.result)(args)
328}
329
330pub fn check_arity(name: &str, positional: usize, named: &[String]) -> Result<(), Mismatch> {
333 let Some(sig) = signature(name) else {
334 return Ok(());
335 };
336 let ceiling = if sig.variadic {
337 usize::from(MAX_WIDTH)
338 } else {
339 sig.params.len()
340 };
341 for key in named {
342 let Some(at) = sig.params.iter().position(|p| p.name == *key) else {
343 if sig.named().contains(&key.as_str()) {
344 continue;
345 }
346 return Err(Mismatch::new(
347 "grammar.unknown_named_argument",
348 &[],
349 format!("`{name}` does not read `{key}`"),
350 ));
351 };
352 let param = sig.params[at];
353 if param.kind == ParamKind::Signal {
354 return Err(Mismatch::new(
355 "grammar.arity",
356 &[],
357 format!("give the signal `{name}` reads by position, not as `{key}=`"),
358 ));
359 }
360 if at < positional {
361 return Err(Mismatch::new(
362 "grammar.arity",
363 &[],
364 format!("give `{key}` to `{name}` by position or by name, not both"),
365 ));
366 }
367 }
368 if positional > ceiling {
369 return Err(Mismatch::new(
370 "grammar.arity",
371 &[],
372 format!("`{name}` takes {} arguments", sig.params.len()),
373 ));
374 }
375 match sig.params[positional.min(sig.params.len())..]
376 .iter()
377 .find(|p| p.required && !named.iter().any(|k| k == p.name))
378 {
379 None => Ok(()),
380 Some(p) if name == "rand" => Err(Mismatch::new(
381 "grammar.arity",
382 &[],
383 format!(
384 "a draw reads its `{}`, and a random source is a function of time: write \
385 `rand(t, seed=k)`, or a constant key for one fixed draw",
386 p.name
387 ),
388 )),
389 Some(p) => Err(Mismatch::new(
390 "grammar.arity",
391 &[],
392 format!("`{name}` reads `{}`, and none was written", p.name),
393 )),
394 }
395}
396
397pub fn notation(ty: Ty) -> &'static str {
399 match (ty.held, ty.has_dual()) {
400 (Held::Form(Var::T), false) => "Form(t)",
401 (Held::Form(Var::T), true) => "Form(t) dual",
402 (Held::Form(Var::F), false) => "Form(f)",
403 (Held::Form(Var::F), true) => "Form(f) dual",
404 (Held::Sampled, _) => "Sampled",
405 (Held::Frames, _) => "Frames",
406 }
407}
408
409pub fn describe(ty: Ty) -> String {
411 let word = match ty.held {
412 Held::Form(var) => format!(
413 "a closed form in `{}` with {} dual",
414 var.as_str(),
415 match ty.has_dual() {
416 true => "a",
417 false => "no",
418 }
419 ),
420 Held::Sampled => "samples".to_string(),
421 Held::Frames => "frames".to_string(),
422 };
423 match (ty.held, ty.rate, ty.width) {
424 (Held::Sampled, Some(rate), _) => format!("samples at {rate} Hz"),
425 (_, _, 1) => word,
426 (_, _, width) => format!("{word}, {width} components"),
427 }
428}
429
430pub fn format_refusal(
433 call: &str,
434 m: &Mismatch,
435 at: Located,
436 operands: &[(String, Ty)],
437) -> Diagnostic {
438 let types: Vec<String> = m.got.iter().map(|t| describe(*t)).collect();
439 let named: Vec<String> = operands
440 .iter()
441 .map(|(name, ty)| format!("{name}: {}", describe(*ty)))
442 .collect();
443 let head = match &m.blocker {
444 Some((_, what)) => {
445 format!("`{call}` refused: subterm {what} left A")
446 }
447 None => format!("`{call}` has no overload for ({}).", types.join(", ")),
448 };
449 let message = match named.is_empty() {
450 true => head,
451 false => format!("{head} {}.", named.join(". ")),
452 };
453 Diagnostic {
454 code: m.code.to_string(),
455 message,
456 location: at,
457 help: m.repair.clone(),
458 }
459}