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