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