1mod draft;
4mod folds;
5
6use std::collections::{BTreeMap, BTreeSet};
7use std::sync::Arc;
8
9use sva_formula::filter::Shape;
10use sva_formula::{ClosedForm, Codomain, Env, Held, NodeId, Origin, ParamId, Ty, Var, infer};
11use sva_samples::Params;
12
13use crate::arguments::{Arguments, Called, Chosen};
14use crate::cast::Cast;
15use crate::error::{Diagnostic, EngineError, Located};
16use crate::instantiate::Instances;
17use crate::lower;
18use crate::schedule::Order;
19use crate::time::Grid;
20
21use draft::{Draft, Entry, Units};
22
23#[derive(Clone, Debug, PartialEq)]
25pub enum Value {
26 ClosedForm(ClosedForm),
27 Cast(Cast, NodeId),
28 Op {
29 name: String,
30 args: Vec<NodeId>,
31 },
32 SelfAt {
33 at: When,
34 },
35 Read {
36 source: NodeId,
37 at: When,
38 site: Origin,
39 },
40 Noise(u64),
41 Filter {
43 shape: Shape,
44 x: NodeId,
45 cutoff: NodeId,
46 q: NodeId,
47 gain: NodeId,
48 },
49 Solver {
50 params: Box<Params>,
51 varying: Vec<(&'static str, NodeId)>,
52 },
53}
54
55#[derive(Clone, Debug, PartialEq)]
58pub enum When {
59 At(crate::time::Affine),
60 Moving(NodeId),
61 Index(crate::index::Index),
62 Step(Step),
63}
64
65#[derive(Clone, Debug, PartialEq)]
68pub enum Step {
69 Index(crate::index::Index),
70 Nearest(NodeId, crate::index::Round),
71 Add(Vec<Step>),
72 Neg(Box<Step>),
73 Mul(Vec<Step>),
74}
75
76impl Step {
77 pub(crate) fn times(&self, out: &mut Vec<NodeId>) {
79 match self {
80 Step::Index(_) => {}
81 Step::Nearest(time, _) => out.push(*time),
82 Step::Add(parts) | Step::Mul(parts) => parts.iter().for_each(|p| p.times(out)),
83 Step::Neg(p) => p.times(out),
84 }
85 }
86}
87
88impl When {
89 pub(crate) fn moving(&self) -> Vec<NodeId> {
91 let mut out = Vec::new();
92 match self {
93 When::Moving(id) => out.push(*id),
94 When::Step(step) => step.times(&mut out),
95 When::At(_) | When::Index(_) => {}
96 }
97 out
98 }
99}
100
101#[derive(Clone, Debug, PartialEq)]
102pub struct Node {
103 pub name: Arc<str>,
104 pub ty: Ty,
105 pub var: Var,
106 pub value: Value,
107 pub grid: Grid,
108}
109
110#[derive(Clone, Debug, Default, PartialEq)]
112pub struct Typing {
113 nodes: Vec<Option<Node>>,
114 free: Vec<u32>,
115 arguments: BTreeMap<String, Arguments>,
116 by_path: BTreeMap<String, NodeId>,
117 aliases: BTreeSet<String>,
119 copies: BTreeMap<String, BTreeMap<Grid, NodeId>>,
120 lowering: BTreeSet<String>,
121 files: BTreeMap<String, Vec<NodeId>>,
122 origins: Vec<Option<(u32, Option<sva_ast::ByteSpan>)>>,
123 free_origins: Vec<u32>,
124 sites: Vec<Option<(String, u32)>>,
126 free_sites: Vec<u32>,
127 site_ids: BTreeMap<String, u32>,
128 pending: BTreeSet<NodeId>,
129 indices: u32,
130 sum: Option<(NodeId, Vec<SumSlot>)>,
131 folds: folds::Folds,
132 readers: BTreeMap<NodeId, BTreeSet<NodeId>>,
134 lowered: Vec<String>,
136 units: BTreeMap<String, Units>,
137 making: Vec<(Arc<str>, Grid)>,
138 draft: Draft,
139}
140
141#[derive(Clone, Copy, Debug, PartialEq)]
143pub(crate) enum SumSlot {
144 Node(NodeId),
145 Retired(sva_samples::Extent),
146}
147
148impl Typing {
149 pub(crate) fn name_sum(&mut self, node: NodeId, slots: Vec<SumSlot>) {
151 self.draft.sum = Some(Some((node, slots)));
152 self.forget([node]);
153 }
154
155 pub(crate) fn forget(&mut self, changed: impl IntoIterator<Item = NodeId>) {
157 let mut gone: BTreeSet<NodeId> = changed.into_iter().collect();
158 if gone.is_empty() {
159 return;
160 }
161 let mut open: Vec<NodeId> = gone.iter().copied().collect();
162 while let Some(at) = open.pop() {
163 let summed = self
164 .summed()
165 .filter(|(_, slots)| slots.contains(&SumSlot::Node(at)));
166 let held = self.readers.get(&at).into_iter().flatten().copied();
167 for reader in held.chain(summed.map(|(sum, _)| *sum)) {
168 if gone.insert(reader) {
169 open.push(reader);
170 }
171 }
172 }
173 self.folds.forget(&gone);
174 }
175
176 pub(super) fn place(&mut self, id: NodeId, node: Option<Node>) {
178 if self.nodes[id.0 as usize].is_some() {
179 for read in self.own_edges(id) {
180 if let Some(held) = self.readers.get_mut(&read)
181 && held.remove(&id)
182 && held.is_empty()
183 {
184 self.readers.remove(&read);
185 }
186 }
187 }
188 self.nodes[id.0 as usize] = node;
189 if self.nodes[id.0 as usize].is_some() {
190 for read in self.own_edges(id) {
191 self.readers.entry(read).or_default().insert(id);
192 }
193 }
194 }
195
196 pub(crate) fn readers_of(&self, id: NodeId) -> impl Iterator<Item = NodeId> + '_ {
198 self.readers.get(&id).into_iter().flatten().copied()
199 }
200
201 fn summed(&self) -> Option<&(NodeId, Vec<SumSlot>)> {
202 match &self.draft.sum {
203 Some(staged) => staged.as_ref(),
204 None => self.sum.as_ref(),
205 }
206 }
207
208 pub(crate) fn retired_sum(&self) -> Option<NodeId> {
210 let (sum, slots) = self.summed()?;
211 let retired = slots.iter().any(|slot| matches!(slot, SumSlot::Retired(_)));
212 retired.then_some(*sum)
213 }
214
215 pub(crate) fn operands(&self, id: NodeId) -> Vec<NodeId> {
217 match self.value(id) {
218 Value::ClosedForm(form) => crate::refs::nodes_in(&form.body),
219 Value::Cast(_, source) | Value::Read { source, .. } => vec![*source],
220 Value::Op { args, .. } => args.clone(),
221 Value::Filter {
222 x, cutoff, q, gain, ..
223 } => vec![*x, *cutoff, *q, *gain],
224 Value::Solver { varying, .. } => varying.iter().map(|(_, a)| *a).collect(),
225 Value::SelfAt { .. } | Value::Noise(_) => Vec::new(),
226 }
227 }
228
229 pub(crate) fn reads(&self, id: NodeId, of: NodeId) -> bool {
231 let (mut open, mut seen) = (vec![id], BTreeSet::new());
232 while let Some(at) = open.pop() {
233 if at == of {
234 return true;
235 }
236 if seen.insert(at) {
237 open.extend(self.edges(at));
238 }
239 }
240 false
241 }
242
243 pub(crate) fn unfolded(&self, root: NodeId, held: impl Fn(NodeId) -> bool) -> Vec<NodeId> {
246 self.unfolded_over(root, |id| self.edges(id), held)
247 }
248
249 pub(crate) fn unfolded_over(
250 &self,
251 root: NodeId,
252 reads: impl Fn(NodeId) -> Vec<NodeId>,
253 held: impl Fn(NodeId) -> bool,
254 ) -> Vec<NodeId> {
255 let mut order = Vec::new();
256 let mut seen = BTreeSet::new();
257 let mut open = vec![(root, false)];
258 while let Some((at, read)) = open.pop() {
259 if read {
260 order.push(at);
261 continue;
262 }
263 if held(at) || !seen.insert(at) {
264 continue;
265 }
266 open.push((at, true));
267 open.extend(reads(at).into_iter().rev().map(|n| (n, false)));
268 }
269 order
270 }
271
272 fn own_edges(&self, id: NodeId) -> Vec<NodeId> {
274 let mut out = self.operands(id);
275 if let Value::Read { at, .. } | Value::SelfAt { at, .. } = self.value(id) {
276 out.extend(at.moving());
277 }
278 out
279 }
280
281 fn edges(&self, id: NodeId) -> Vec<NodeId> {
283 let mut out = self.own_edges(id);
284 let slots = self.sum_slots(id).unwrap_or_default().iter();
285 out.extend(slots.filter_map(|slot| match slot {
286 SumSlot::Node(read) if *read != id => Some(*read),
287 _ => None,
288 }));
289 out
290 }
291
292 pub(crate) fn sum_slots(&self, node: NodeId) -> Option<&[SumSlot]> {
293 self.summed()
294 .filter(|(held, _)| *held == node)
295 .map(|(_, slots)| slots.as_slice())
296 }
297
298 pub(crate) fn lowering(&mut self, path: &str) {
299 self.lowered.push(path.to_string());
300 }
301
302 pub(crate) fn lowered(&self) -> &[String] {
303 &self.lowered
304 }
305
306 pub(crate) fn next_index(&mut self) -> sva_formula::IndexId {
307 self.indices += 1;
308 sva_formula::IndexId(self.indices)
309 }
310
311 pub(crate) fn note(&mut self, node: &str, call: Option<Called>, chosen: Vec<Chosen>) {
313 let mut held = self.arguments(node).cloned().unwrap_or_else(|| Arguments {
314 node: node.to_string(),
315 ..Arguments::default()
316 });
317 if let Some(call) = call {
318 held.calls
319 .retain(|c| (c.at.start, &c.name) != (call.at.start, &call.name));
320 held.calls.push(call);
321 }
322 for one in chosen {
323 held.chosen.retain(|c| c.at.start != one.at.start);
324 held.chosen.push(one);
325 }
326 let old = self.draft.arguments.insert(node.to_string(), held);
327 self.draft.journal.push(Entry::Noted(node.to_string(), old));
328 }
329
330 pub fn arguments(&self, node: &str) -> Option<&Arguments> {
331 match self.draft.arguments.get(node) {
332 Some(held) => Some(held),
333 None if self.draft.hidden.contains(node) => None,
334 None => self.arguments.get(node),
335 }
336 }
337
338 pub fn ty(&self, n: NodeId) -> Ty {
339 self.at(n).ty
340 }
341
342 pub fn var(&self, n: NodeId) -> Var {
343 self.at(n).var
344 }
345
346 pub fn value(&self, n: NodeId) -> &Value {
347 &self.at(n).value
348 }
349
350 pub fn name(&self, n: NodeId) -> &str {
351 &self.at(n).name
352 }
353
354 pub fn grid(&self, n: NodeId) -> Grid {
355 self.at(n).grid
356 }
357
358 pub(crate) fn copy(&self, path: &str, grid: Grid) -> Option<NodeId> {
359 let key = (path.to_string(), grid);
360 match self.draft.copies.get(&key) {
361 Some(id) => Some(*id),
362 None if self.draft.hidden.contains(path) => None,
363 None => self
364 .copies
365 .get(path)
366 .and_then(|held| held.get(&grid))
367 .copied(),
368 }
369 }
370
371 pub(crate) fn copied(&mut self, path: &str, grid: Grid, id: NodeId) {
372 let key = (path.to_string(), grid);
373 let old = self.draft.copies.insert(key.clone(), id);
374 self.draft.journal.push(Entry::Copy(key, old));
375 }
376
377 pub(crate) fn opened(&mut self, path: &str) -> bool {
379 self.lowering.insert(path.to_string())
380 }
381
382 pub(crate) fn closed(&mut self, path: &str) {
383 self.lowering.remove(path);
384 }
385
386 pub fn at(&self, n: NodeId) -> &Node {
387 self.nodes[n.0 as usize]
388 .as_ref()
389 .expect("a node the typing holds")
390 }
391
392 pub fn ids(&self) -> impl Iterator<Item = NodeId> + '_ {
393 let held = self.nodes.iter().enumerate();
394 held.filter(|(_, n)| n.is_some())
395 .map(|(at, _)| NodeId(at as u32))
396 }
397
398 #[cfg(test)]
399 pub(crate) fn len(&self) -> usize {
400 self.nodes.len() - self.free.len()
401 }
402
403 pub fn id(&self, path: &str) -> Option<NodeId> {
404 match self.draft.by_path.get(path) {
405 Some(id) => Some(*id),
406 None if self.draft.hidden.contains(path) => None,
407 None => self.by_path.get(path).copied(),
408 }
409 }
410
411 pub fn paths(&self) -> impl Iterator<Item = (&str, NodeId)> {
413 let drafted = self.draft.by_path.iter();
414 let kept = self.by_path.iter().filter(|(path, _)| {
415 !self.draft.hidden.contains(*path) && !self.draft.by_path.contains_key(*path)
416 });
417 let all: BTreeMap<&String, &NodeId> = drafted.chain(kept).collect();
418 all.into_iter().map(|(p, id)| (p.as_str(), *id))
419 }
420
421 pub fn resolve(&self, path: &str) -> Result<NodeId, EngineError> {
423 if let Some(id) = self.id(path) {
424 return Ok(id);
425 }
426 let held = self.files.get(path).cloned().unwrap_or_default();
427 crate::instantiate::sole(
428 path,
429 held.into_iter()
430 .map(|id| (self.name(id).to_string(), id))
431 .collect(),
432 )
433 }
434
435 pub fn locate(&self, origin: Origin) -> Located {
437 let held = self.origins.get(origin.token() as usize).copied().flatten();
438 held.and_then(|(site, span)| {
439 let (name, _) = self.sites[site as usize].as_ref()?;
440 Some(Located::at(name.as_str(), span))
441 })
442 .unwrap_or_default()
443 }
444
445 pub(crate) fn mark(&mut self, node: &str, span: Option<sva_ast::ByteSpan>) -> Origin {
446 let site = match self.site_ids.get(node) {
447 Some(&site) => site,
448 None => {
449 let held = Some((node.to_string(), 0));
450 let site = match self.free_sites.pop() {
451 Some(site) => {
452 self.sites[site as usize] = held;
453 site
454 }
455 None => {
456 self.sites.push(held);
457 (self.sites.len() - 1) as u32
458 }
459 };
460 self.site_ids.insert(node.to_string(), site);
461 site
462 }
463 };
464 self.sites[site as usize].as_mut().expect("a site").1 += 1;
465 let token = match self.free_origins.pop() {
466 Some(token) => {
467 self.origins[token as usize] = Some((site, span));
468 token
469 }
470 None => {
471 self.origins.push(Some((site, span)));
472 (self.origins.len() - 1) as u32
473 }
474 };
475 let unit = self.unit();
476 self.draft.journal.push(Entry::Origin(token, unit));
477 Origin::new(token)
478 }
479
480 pub(crate) fn seed(&mut self, path: &str, held: Held, grid: Grid) -> NodeId {
481 if let Some(id) = self.id(path) {
482 return id;
483 }
484 let name = self.begin(path, grid);
485 let id = self.push(
486 Node {
487 name,
488 ty: Ty {
489 dual: held.is_closed_form(),
490 ..Ty::discrete(held, Codomain::Real)
491 },
492 var: Var::T,
493 value: Value::Op {
494 name: "loop".to_string(),
495 args: Vec::new(),
496 },
497 grid,
498 },
499 Some(path),
500 );
501 self.end();
502 self.pending.insert(id);
503 self.draft.journal.push(Entry::Pending(id));
504 id
505 }
506
507 pub(crate) fn pending(&self, id: NodeId) -> bool {
508 self.pending.contains(&id)
509 }
510
511 pub(crate) fn settle(&mut self, id: NodeId, node: Node) {
513 self.place(id, Some(node));
514 self.pending.remove(&id);
515 self.forget([id]);
516 }
517
518 pub(crate) fn folds(&self) -> &folds::Folds {
519 &self.folds
520 }
521
522 pub(crate) fn alias(&mut self, path: &str, id: NodeId) {
523 let old = self.draft.by_path.insert(path.to_string(), id);
524 self.draft.journal.push(Entry::Path(path.to_string(), old));
525 }
526
527 pub(crate) fn push(&mut self, node: Node, path: Option<&str>) -> NodeId {
528 let id = match self.free.pop() {
529 Some(at) => NodeId(at),
530 None => {
531 self.nodes.push(None);
532 NodeId((self.nodes.len() - 1) as u32)
533 }
534 };
535 self.place(id, Some(node));
536 let unit = self.unit();
537 self.draft.journal.push(Entry::Node(id, unit));
538 if let Some(path) = path {
539 self.alias(path, id);
540 }
541 id
542 }
543
544 pub(crate) fn begin(&mut self, path: &str, grid: Grid) -> Arc<str> {
547 let name: Arc<str> = path.into();
548 self.making.push((Arc::clone(&name), grid));
549 name
550 }
551
552 pub(crate) fn end(&mut self) {
553 self.making.pop();
554 }
555
556 fn unit(&self) -> (Arc<str>, Grid) {
557 self.making
558 .last()
559 .cloned()
560 .unwrap_or_else(|| (Arc::from(""), Grid::of(1)))
561 }
562
563 pub(crate) fn infer_closed_form(&self, form: &ClosedForm) -> Result<Ty, EngineError> {
564 let inferred = crate::refs::read_through(self, |through| {
565 infer(form, &TypedNodesEnv(&self.nodes, through))
566 });
567 inferred.map_err(|r| {
568 EngineError::of_closed_form(
569 &r,
570 self.locate(r.origin),
571 "write the subterm inside sample(...) to leave A deliberately",
572 )
573 })
574 }
575}
576
577struct TypedNodesEnv<'a>(&'a [Option<Node>], &'a dyn sva_formula::Reads);
578
579impl Env for TypedNodesEnv<'_> {
580 fn node(&self, id: NodeId) -> Ty {
581 self.0[id.0 as usize]
582 .as_ref()
583 .expect("a node the typing holds")
584 .ty
585 }
586
587 fn param(&self, _: ParamId) -> Ty {
589 Ty::form(Var::T, true, Codomain::Real)
590 }
591
592 fn reads(&self) -> &dyn sva_formula::Reads {
593 self.1
594 }
595}
596
597pub fn infer_all(inst: &Instances, order: &Order) -> Result<Typing, EngineError> {
599 let mut typing = Typing::default();
600 typing.lower(inst, &order.groups)?;
601 typing.commit(inst);
602 Ok(typing)
603}
604
605impl Typing {
606 pub(crate) fn lower(
608 &mut self,
609 inst: &Instances,
610 groups: &[Vec<String>],
611 ) -> Result<(), EngineError> {
612 self.lowered.clear();
613 self.hide(groups.iter().flatten().cloned());
614 for group in groups {
615 match crate::schedule::is_loop(inst, group) {
616 true => settle_loop(self, inst, group)?,
617 false => {
618 for path in group {
619 lower::node(path, inst, self)?;
620 }
621 }
622 }
623 }
624 Ok(())
625 }
626}
627
628const SEEDS: [Held; 2] = [Held::Form(Var::T), Held::Sampled];
629
630fn settle_loop(typing: &mut Typing, inst: &Instances, group: &[String]) -> Result<(), EngineError> {
633 let mut seed = SEEDS[0];
634 let mut refusal = None;
635 for _ in 0..=SEEDS.len() {
636 let mark = typing.checkpoint();
637 for path in group {
638 typing.seed(path, seed, inst.grid());
639 }
640 let walked = group
641 .iter()
642 .try_for_each(|path| lower::node(path, inst, typing).map(|_| ()));
643 match walked {
644 Ok(()) => match held_by(typing, group) {
645 reached if reached == seed => return Ok(()),
646 reached => seed = reached,
647 },
648 Err(e) => {
649 refusal = Some(e);
650 seed = elsewhere(seed);
651 }
652 }
653 typing.rollback(mark);
654 }
655 Err(refusal.unwrap_or_else(|| mixed(group)))
656}
657
658fn elsewhere(seed: Held) -> Held {
660 match seed {
661 Held::Sampled => Held::Form(Var::T),
662 _ => Held::Sampled,
663 }
664}
665
666fn held_by(typing: &Typing, group: &[String]) -> Held {
668 let sampled = group
669 .iter()
670 .filter_map(|path| typing.id(path))
671 .any(|id| !typing.ty(id).is_closed_form());
672 match sampled {
673 true => Held::Sampled,
674 false => Held::Form(Var::T),
675 }
676}
677
678fn mixed(group: &[String]) -> EngineError {
680 let at = group.first().map_or("", String::as_str);
681 EngineError::refused(Diagnostic {
682 code: "type.samples_in_closed_form".to_string(),
683 message: format!(
684 "the loop over `{}` holds a closed form and samples at once.",
685 group.join("`, `")
686 ),
687 location: Located::at(at, None),
688 help: "write sample(...) on the members that are closed forms, so the whole loop runs \
689 on the grid"
690 .to_string(),
691 })
692}