1use std::sync::Arc;
2
3use sim_kernel::{
4 AbiVersion, Args, Callable, ClassRef, Cx, Error, Export, ExportKind, ExportRecord, ExportState,
5 Expr, Lib, LibManifest, LibTarget, Linker, NumberLiteral, Object, ObjectCompat, RawArgs,
6 Result, RuntimeId, ShapeRef, Symbol, Value, Version,
7};
8use sim_lib_music_core::Time;
9use sim_lib_music_shapes::{decode_counterpoint, decode_melody, install_music_shapes_lib};
10use sim_shape::{AnyShape, ExactExprShape, ListShape, shape_value};
11
12use crate::runtime_generation::generate_call;
13use crate::runtime_graph_expr::stretto_graph_expr;
14use crate::{
15 CounterpointReport, MetricEvidence, NoteEvidence, RuleSet, StrettoPolicy, StrettoTransform,
16 TimeSpan, Violation, VoiceEvidence, analyze_counterpoint, stretto_graph,
17};
18
19const LIB_ID: &str = "music-counterpoint";
20const EXPORT_KIND: &str = "CounterpointAnalyzer";
21
22pub struct MusicCounterpointLib;
24
25impl Lib for MusicCounterpointLib {
26 fn manifest(&self) -> LibManifest {
27 LibManifest {
28 id: Symbol::new(LIB_ID),
29 version: Version(env!("CARGO_PKG_VERSION").to_owned()),
30 abi: AbiVersion { major: 0, minor: 1 },
31 target: LibTarget::HostRegistered,
32 requires: Vec::new(),
33 capabilities: Vec::new(),
34 exports: vec![
35 Export::Value {
36 symbol: analyzer_symbol(),
37 },
38 Export::Function {
39 symbol: music_counterpoint_analyze_symbol(),
40 function_id: None,
41 },
42 Export::Function {
43 symbol: music_counterpoint_generate_symbol(),
44 function_id: None,
45 },
46 Export::Function {
47 symbol: music_stretto_graph_symbol(),
48 function_id: None,
49 },
50 ],
51 }
52 }
53
54 fn load(&self, cx: &mut sim_kernel::LoadCx, linker: &mut Linker<'_>) -> Result<()> {
55 linker.value(analyzer_symbol(), analyzer_value(cx)?)?;
56 linker.function_value(
57 music_counterpoint_analyze_symbol(),
58 cx.factory()
59 .opaque(Arc::new(CounterpointFunction::Analyze))?,
60 )?;
61 linker.function_value(
62 music_counterpoint_generate_symbol(),
63 cx.factory()
64 .opaque(Arc::new(CounterpointFunction::Generate))?,
65 )?;
66 linker.function_value(
67 music_stretto_graph_symbol(),
68 cx.factory()
69 .opaque(Arc::new(CounterpointFunction::Stretto))?,
70 )?;
71 Ok(())
72 }
73}
74
75pub fn install_music_counterpoint_lib(cx: &mut Cx) -> Result<()> {
77 install_music_shapes_lib(cx)?;
78 if !sim_lib_core::install_once(cx, &MusicCounterpointLib)? {
79 return Ok(());
80 }
81 cx.registry_mut().append_export_record(
82 &Symbol::new(LIB_ID),
83 ExportRecord {
84 kind: ExportKind::named(EXPORT_KIND),
85 symbol: analyzer_symbol(),
86 state: ExportState::Resolved {
87 id: RuntimeId::Value,
88 },
89 },
90 )?;
91 Ok(())
92}
93
94pub fn music_counterpoint_analyze_symbol() -> Symbol {
96 Symbol::qualified("music/counterpoint", "analyze")
97}
98
99pub fn music_stretto_graph_symbol() -> Symbol {
101 Symbol::qualified("music/stretto", "graph")
102}
103
104pub fn music_counterpoint_generate_symbol() -> Symbol {
106 Symbol::qualified("music/counterpoint", "generate")
107}
108
109fn analyzer_symbol() -> Symbol {
110 Symbol::qualified("music", "CounterpointAnalyzer")
111}
112
113fn analyzer_value(cx: &mut sim_kernel::LoadCx) -> Result<Value> {
114 cx.factory().table(vec![
115 (
116 Symbol::new("symbol"),
117 cx.factory().symbol(analyzer_symbol())?,
118 ),
119 (
120 Symbol::new("layer"),
121 cx.factory().string("music".to_owned())?,
122 ),
123 (
124 Symbol::new("kind"),
125 cx.factory().string("analysis-and-generation".to_owned())?,
126 ),
127 (
128 Symbol::new("shape"),
129 cx.factory()
130 .symbol(Symbol::qualified("music", "CounterpointAnalyzer"))?,
131 ),
132 (
133 Symbol::new("dependencies"),
134 cx.factory().list(
135 [
136 "music-core",
137 "music-consonance",
138 "music-transform",
139 "discrete-graph",
140 "discrete-search",
141 ]
142 .into_iter()
143 .map(|name| cx.factory().string(name.to_owned()))
144 .collect::<Result<Vec<_>>>()?,
145 )?,
146 ),
147 (Symbol::new("lossless"), cx.factory().bool(true)?),
148 (Symbol::new("capabilities"), cx.factory().list(Vec::new())?),
149 (
150 Symbol::new("analysis-callable"),
151 cx.factory().symbol(music_counterpoint_analyze_symbol())?,
152 ),
153 (
154 Symbol::new("stretto-callable"),
155 cx.factory().symbol(music_stretto_graph_symbol())?,
156 ),
157 (
158 Symbol::new("generation-callable"),
159 cx.factory().symbol(music_counterpoint_generate_symbol())?,
160 ),
161 (Symbol::new("generation"), cx.factory().bool(true)?),
162 ])
163}
164
165enum CounterpointFunction {
166 Analyze,
167 Generate,
168 Stretto,
169}
170
171impl Object for CounterpointFunction {
172 fn display(&self, _cx: &mut Cx) -> Result<String> {
173 Ok(match self {
174 Self::Analyze => "#<function music/counterpoint/analyze>",
175 Self::Generate => "#<function music/counterpoint/generate>",
176 Self::Stretto => "#<function music/stretto/graph>",
177 }
178 .to_owned())
179 }
180
181 fn as_any(&self) -> &dyn std::any::Any {
182 self
183 }
184}
185
186impl ObjectCompat for CounterpointFunction {
187 fn class(&self, cx: &mut Cx) -> Result<ClassRef> {
188 cx.factory().class_stub(
189 sim_kernel::CORE_FUNCTION_CLASS_ID,
190 Symbol::qualified("core", "Function"),
191 )
192 }
193
194 fn as_callable(&self) -> Option<&dyn Callable> {
195 Some(self)
196 }
197}
198
199impl Callable for CounterpointFunction {
200 fn call(&self, cx: &mut Cx, args: Args) -> Result<Value> {
201 let exprs = args
202 .into_vec()
203 .into_iter()
204 .map(|value| value.object().as_expr(cx))
205 .collect::<Result<Vec<_>>>()?;
206 self.invoke(cx, &exprs, false)
207 }
208
209 fn call_exprs(&self, cx: &mut Cx, args: RawArgs) -> Result<Value> {
210 self.invoke(cx, args.exprs(), true)
211 }
212
213 fn browse_args_shape(&self, _cx: &mut Cx) -> Result<Option<ShapeRef>> {
214 let keyword = |name| {
215 Arc::new(ExactExprShape::new(Expr::Symbol(Symbol::new(name))))
216 as Arc<dyn sim_shape::Shape>
217 };
218 let fields = match self {
219 Self::Analyze => vec![
220 keyword(":counterpoint"),
221 Arc::new(AnyShape),
222 keyword(":rules"),
223 Arc::new(AnyShape),
224 ],
225 Self::Generate => vec![
226 Arc::new(AnyShape),
227 keyword(":rules"),
228 Arc::new(AnyShape),
229 keyword(":voices"),
230 Arc::new(AnyShape),
231 keyword(":control"),
232 Arc::new(AnyShape),
233 ],
234 Self::Stretto => vec![
235 keyword(":subject"),
236 Arc::new(AnyShape),
237 keyword(":policy"),
238 Arc::new(AnyShape),
239 ],
240 };
241 Ok(Some(shape_value(
242 match self {
243 Self::Analyze => Symbol::qualified("music/counterpoint/analyze", "args"),
244 Self::Generate => Symbol::qualified("music/counterpoint/generate", "args"),
245 Self::Stretto => Symbol::qualified("music/stretto/graph", "args"),
246 },
247 Arc::new(ListShape::tuple(fields)),
248 )))
249 }
250
251 fn browse_result_shape(&self, _cx: &mut Cx) -> Result<Option<ShapeRef>> {
252 Ok(Some(shape_value(
253 match self {
254 Self::Analyze => Symbol::qualified("music/counterpoint/analyze", "result"),
255 Self::Generate => Symbol::qualified("music/counterpoint/generate", "result"),
256 Self::Stretto => Symbol::qualified("music/stretto/graph", "result"),
257 },
258 Arc::new(AnyShape),
259 )))
260 }
261}
262
263impl CounterpointFunction {
264 fn invoke(&self, cx: &mut Cx, args: &[Expr], evaluate_values: bool) -> Result<Value> {
265 match self {
266 Self::Analyze => analyze_call(cx, args, evaluate_values),
267 Self::Generate => generate_call(cx, args, evaluate_values),
268 Self::Stretto => stretto_call(cx, args, evaluate_values),
269 }
270 }
271}
272
273fn analyze_call(cx: &mut Cx, args: &[Expr], evaluate_values: bool) -> Result<Value> {
274 let [counterpoint_key, counterpoint, rules_key, rules] = args else {
275 return Err(Error::Eval(
276 "music/counterpoint/analyze expects :counterpoint STRING :rules SYMBOL".to_owned(),
277 ));
278 };
279 expect_keyword(counterpoint_key, "counterpoint")?;
280 expect_keyword(rules_key, "rules")?;
281 let counterpoint = value_expr(cx, counterpoint, evaluate_values)?;
282 let Expr::String(counterpoint) = unquote(counterpoint) else {
283 return Err(Error::TypeMismatch {
284 expected: "canonical #(Counterpoint ...) string",
285 found: "non-string",
286 });
287 };
288 let counterpoint = decode_counterpoint(&counterpoint)
289 .map_err(|error| Error::Eval(format!("invalid counterpoint: {error}")))?;
290 let rules = value_expr(cx, rules, evaluate_values)?;
291 let rules = named_rules(&symbolish(&rules)?)?;
292 cx.factory()
293 .expr(counterpoint_report_expr(&analyze_counterpoint(
294 &counterpoint,
295 &rules,
296 )))
297}
298
299fn stretto_call(cx: &mut Cx, args: &[Expr], evaluate_values: bool) -> Result<Value> {
300 let [subject_key, subject, policy_key, policy] = args else {
301 return Err(Error::Eval(
302 "music/stretto/graph expects :subject STRING :policy MAP".to_owned(),
303 ));
304 };
305 expect_keyword(subject_key, "subject")?;
306 expect_keyword(policy_key, "policy")?;
307 let subject = value_expr(cx, subject, evaluate_values)?;
308 let Expr::String(subject) = unquote(subject) else {
309 return Err(Error::TypeMismatch {
310 expected: "canonical #(Melody ...) string",
311 found: "non-string",
312 });
313 };
314 let subject = decode_melody(&subject)
315 .map_err(|error| Error::Eval(format!("invalid stretto subject: {error}")))?;
316 let policy = value_expr(cx, policy, evaluate_values)?;
317 let policy = parse_stretto_policy(&policy)?;
318 let graph = stretto_graph(&subject, policy).map_err(|error| Error::Eval(error.to_string()))?;
319 cx.factory().expr(stretto_graph_expr(&graph))
320}
321
322pub(crate) fn named_rules(name: &str) -> Result<RuleSet> {
323 let pulse = Time::from_integer(1);
324 match name {
325 "species-one" => Ok(RuleSet::species_one(pulse)),
326 "species-two" => Ok(RuleSet::species_two(pulse)),
327 "species-three" => Ok(RuleSet::species_three(pulse)),
328 "species-four" => Ok(RuleSet::species_four(pulse)),
329 "open" => Ok(RuleSet::open()),
330 other => Err(Error::Eval(format!(
331 "unknown counterpoint rule set {other}"
332 ))),
333 }
334}
335
336fn parse_stretto_policy(expr: &Expr) -> Result<StrettoPolicy> {
337 let Expr::Map(entries) = unquote_ref(expr) else {
338 return Err(Error::TypeMismatch {
339 expected: "stretto policy map",
340 found: "non-map",
341 });
342 };
343 let mut policy = StrettoPolicy::default();
344 for (key, value) in entries {
345 match keyword_name(key)?.as_str() {
346 "delays" => policy.delays = scalar_list(value, parse_time)?,
347 "transpositions" => {
348 policy.transforms = scalar_list(value, parse_i32)?
349 .into_iter()
350 .map(StrettoTransform::original)
351 .collect();
352 }
353 "minimum-overlap" => policy.minimum_overlap = parse_time(value)?,
354 "minimum-cluster-voices" => policy.minimum_cluster_voices = parse_usize(value)?,
355 "max-entries" => policy.max_entries = parse_usize(value)?,
356 "max-clusters" => policy.max_clusters = parse_usize(value)?,
357 "max-chain-length" => policy.max_chain_length = parse_usize(value)?,
358 "rules" => policy.compatibility_rules = named_rules(&symbolish(value)?)?,
359 other => {
360 return Err(Error::Eval(format!(
361 "unknown music/stretto policy :{other}"
362 )));
363 }
364 }
365 }
366 Ok(policy)
367}
368
369fn counterpoint_report_expr(report: &CounterpointReport) -> Expr {
370 map(vec![
371 ("mode", Expr::String(report.provenance.mode.clone())),
372 ("rule-set", Expr::String(report.provenance.rule_set.clone())),
373 ("facts", strings(&report.provenance.facts)),
374 (
375 "alignment",
376 Expr::Vector(
377 report
378 .alignment
379 .iter()
380 .map(|window| {
381 map(vec![
382 ("span", span_expr(&window.span)),
383 (
384 "notes",
385 Expr::Vector(window.notes.iter().map(note_expr).collect()),
386 ),
387 ])
388 })
389 .collect(),
390 ),
391 ),
392 (
393 "motions",
394 Expr::Vector(
395 report
396 .motions
397 .iter()
398 .map(|motion| {
399 map(vec![
400 ("span", span_expr(&motion.span)),
401 (
402 "voices",
403 Expr::Vector(motion.voices.iter().map(voice_expr).collect()),
404 ),
405 (
406 "notes",
407 Expr::Vector(motion.notes.iter().map(note_expr).collect()),
408 ),
409 (
410 "directions",
411 Expr::Vector(vec![
412 symbol(&format!("{:?}", motion.first).to_lowercase()),
413 symbol(&format!("{:?}", motion.second).to_lowercase()),
414 ]),
415 ),
416 ("interval-before", integer(motion.interval_before)),
417 ("interval-after", integer(motion.interval_after)),
418 ])
419 })
420 .collect(),
421 ),
422 ),
423 (
424 "violations",
425 Expr::Vector(report.violations.iter().map(violation_expr).collect()),
426 ),
427 ])
428}
429
430pub(crate) fn violation_expr(violation: &Violation) -> Expr {
431 map(vec![
432 ("rule", Expr::String(violation.rule.clone())),
433 ("message", Expr::String(violation.message.clone())),
434 ("span", span_expr(&violation.span)),
435 (
436 "voices",
437 Expr::Vector(violation.voices.iter().map(voice_expr).collect()),
438 ),
439 (
440 "notes",
441 Expr::Vector(violation.notes.iter().map(note_expr).collect()),
442 ),
443 ("metric", metric_expr(&violation.metric)),
444 ])
445}
446
447fn voice_expr(voice: &VoiceEvidence) -> Expr {
448 map(vec![
449 ("index", integer(voice.index)),
450 ("id", Expr::String(voice.id.to_string())),
451 ("name", Expr::String(voice.name.clone())),
452 ])
453}
454
455fn note_expr(note: &NoteEvidence) -> Expr {
456 map(vec![
457 ("voice", voice_expr(¬e.voice)),
458 ("index", integer(note.index)),
459 ("note-id", Expr::String(note.note_id.to_string())),
460 ("event-id", Expr::String(note.event_id.to_string())),
461 ("span", span_expr(¬e.span)),
462 (
463 "pitch",
464 Expr::String(format!(
465 "{}{}",
466 note.pitch.class.canonical_name(),
467 note.pitch.octave
468 )),
469 ),
470 ])
471}
472
473fn metric_expr(metric: &MetricEvidence) -> Expr {
474 map(vec![
475 ("name", Expr::String(metric.metric.clone())),
476 ("observed", Expr::String(metric.observed.clone())),
477 ("expected", Expr::String(metric.expected.clone())),
478 ("unit", Expr::String(metric.unit.clone())),
479 ("facts", strings(&metric.facts)),
480 ])
481}
482
483pub(crate) fn span_expr(span: &TimeSpan) -> Expr {
484 map(vec![
485 ("start", time_expr(span.start)),
486 ("end", time_expr(span.end)),
487 ])
488}
489
490fn scalar_list<T>(expr: &Expr, parser: impl Fn(&Expr) -> Result<T>) -> Result<Vec<T>> {
491 match unquote_ref(expr) {
492 Expr::List(values) | Expr::Vector(values) => values.iter().map(parser).collect(),
493 _ => Err(Error::TypeMismatch {
494 expected: "list or vector",
495 found: "non-list",
496 }),
497 }
498}
499
500fn parse_time(expr: &Expr) -> Result<Time> {
501 let value = scalar_text(expr)?;
502 let (numerator, denominator) = value.split_once('/').unwrap_or((&value, "1"));
503 let numerator = numerator
504 .parse::<i64>()
505 .map_err(|_| Error::Eval(format!("invalid rational time {value}")))?;
506 let denominator = denominator
507 .parse::<i64>()
508 .map_err(|_| Error::Eval(format!("invalid rational time {value}")))?;
509 if denominator == 0 {
510 return Err(Error::Eval("rational time denominator is zero".to_owned()));
511 }
512 Ok(Time::new(numerator, denominator))
513}
514
515fn parse_i32(expr: &Expr) -> Result<i32> {
516 let value = scalar_text(expr)?;
517 value
518 .parse()
519 .map_err(|_| Error::Eval(format!("invalid i32 {value}")))
520}
521
522pub(crate) fn parse_usize(expr: &Expr) -> Result<usize> {
523 let value = scalar_text(expr)?;
524 value
525 .parse()
526 .map_err(|_| Error::Eval(format!("invalid usize {value}")))
527}
528
529pub(crate) fn scalar_text(expr: &Expr) -> Result<String> {
530 match unquote_ref(expr) {
531 Expr::String(value) => Ok(value.clone()),
532 Expr::Symbol(value) => Ok(value.name.to_string()),
533 Expr::Number(value) => Ok(value.canonical.clone()),
534 _ => Err(Error::TypeMismatch {
535 expected: "string, symbol, or number",
536 found: "compound expression",
537 }),
538 }
539}
540
541pub(crate) fn value_expr(cx: &mut Cx, expr: &Expr, evaluate: bool) -> Result<Expr> {
542 if evaluate {
543 cx.eval_expr(expr.clone())?.object().as_expr(cx)
544 } else {
545 Ok(expr.clone())
546 }
547}
548
549pub(crate) fn expect_keyword(expr: &Expr, expected: &str) -> Result<()> {
550 if keyword_name(expr)? == expected {
551 Ok(())
552 } else {
553 Err(Error::Eval(format!("expected :{expected}")))
554 }
555}
556
557pub(crate) fn keyword_name(expr: &Expr) -> Result<String> {
558 match unquote_ref(expr) {
559 Expr::Symbol(symbol) => Ok(symbol
560 .name
561 .strip_prefix(':')
562 .unwrap_or(symbol.name.as_ref())
563 .to_owned()),
564 _ => Err(Error::TypeMismatch {
565 expected: "keyword symbol",
566 found: "non-symbol",
567 }),
568 }
569}
570
571pub(crate) fn symbolish(expr: &Expr) -> Result<String> {
572 match unquote_ref(expr) {
573 Expr::Symbol(value) => Ok(value.name.to_string()),
574 Expr::String(value) => Ok(value.clone()),
575 _ => Err(Error::TypeMismatch {
576 expected: "symbol or string",
577 found: "other expression",
578 }),
579 }
580}
581
582pub(crate) fn unquote(expr: Expr) -> Expr {
583 match expr {
584 Expr::Quote { expr, .. } => *expr,
585 other => other,
586 }
587}
588
589pub(crate) fn unquote_ref(expr: &Expr) -> &Expr {
590 match expr {
591 Expr::Quote { expr, .. } => expr,
592 other => other,
593 }
594}
595
596pub(crate) fn strings(values: &[String]) -> Expr {
597 Expr::Vector(values.iter().cloned().map(Expr::String).collect())
598}
599
600pub(crate) fn integer(value: impl ToString) -> Expr {
601 Expr::Number(NumberLiteral {
602 domain: Symbol::qualified("numbers", "i64"),
603 canonical: value.to_string(),
604 })
605}
606
607pub(crate) fn time_expr(value: Time) -> Expr {
608 Expr::String(format!("{}/{}", value.numer(), value.denom()))
609}
610
611pub(crate) fn map(entries: Vec<(&str, Expr)>) -> Expr {
612 Expr::Map(
613 entries
614 .into_iter()
615 .map(|(key, value)| (symbol(key), value))
616 .collect(),
617 )
618}
619
620pub(crate) fn symbol(value: &str) -> Expr {
621 Expr::Symbol(Symbol::new(value))
622}