1mod draft;
4
5use std::cell::RefCell;
6use std::collections::{BTreeMap, BTreeSet};
7use std::sync::Arc;
8
9use sva_formula::filter::Shape;
10use sva_formula::{C64, 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(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 numbers: Numbers,
135 lowered: Vec<String>,
137 units: BTreeMap<String, Units>,
138 making: Vec<(String, Grid)>,
139 draft: Draft,
140}
141
142#[derive(Debug, Default)]
144struct Numbers(RefCell<BTreeMap<NodeId, Option<C64>>>);
145
146impl PartialEq for Numbers {
147 fn eq(&self, _: &Numbers) -> bool {
148 true
149 }
150}
151
152#[derive(Clone, Copy, Debug, PartialEq)]
154pub(crate) enum SumSlot {
155 Node(NodeId),
156 Retired(sva_samples::Extent),
157}
158
159impl Typing {
160 pub(crate) fn name_sum(&mut self, node: NodeId, slots: Vec<SumSlot>) {
162 self.draft.sum = Some(Some((node, slots)));
163 }
164
165 fn summed(&self) -> Option<&(NodeId, Vec<SumSlot>)> {
166 match &self.draft.sum {
167 Some(staged) => staged.as_ref(),
168 None => self.sum.as_ref(),
169 }
170 }
171
172 pub(crate) fn retired_sum(&self) -> Option<NodeId> {
174 let (sum, slots) = self.summed()?;
175 let retired = slots.iter().any(|slot| matches!(slot, SumSlot::Retired(_)));
176 retired.then_some(*sum)
177 }
178
179 pub(crate) fn operands(&self, id: NodeId) -> Vec<NodeId> {
181 match self.value(id) {
182 Value::ClosedForm(form) => crate::refs::nodes_in(&form.body),
183 Value::Cast(_, source) | Value::Read { source, .. } => vec![*source],
184 Value::Op { args, .. } => args.clone(),
185 Value::Filter {
186 x, cutoff, q, gain, ..
187 } => vec![*x, *cutoff, *q, *gain],
188 Value::Solver { varying, .. } => varying.iter().map(|(_, a)| *a).collect(),
189 Value::SelfAt { .. } | Value::Noise(_) | Value::Stored(_) => Vec::new(),
190 }
191 }
192
193 pub(crate) fn reads(&self, id: NodeId, of: NodeId) -> bool {
195 let (mut open, mut seen) = (vec![id], BTreeSet::new());
196 while let Some(at) = open.pop() {
197 if at == of {
198 return true;
199 }
200 if seen.insert(at) {
201 open.extend(self.operands(at));
202 }
203 }
204 false
205 }
206
207 pub(crate) fn sum_slots(&self, node: NodeId) -> Option<&[SumSlot]> {
208 self.summed()
209 .filter(|(held, _)| *held == node)
210 .map(|(_, slots)| slots.as_slice())
211 }
212
213 pub(crate) fn lowering(&mut self, path: &str) {
214 self.lowered.push(path.to_string());
215 }
216
217 pub(crate) fn lowered(&self) -> &[String] {
218 &self.lowered
219 }
220
221 pub(crate) fn next_index(&mut self) -> sva_formula::IndexId {
222 self.indices += 1;
223 sva_formula::IndexId(self.indices)
224 }
225
226 pub(crate) fn note(&mut self, node: &str, call: Option<Called>, chosen: Vec<Chosen>) {
228 let mut held = self.arguments(node).cloned().unwrap_or_else(|| Arguments {
229 node: node.to_string(),
230 ..Arguments::default()
231 });
232 if let Some(call) = call {
233 held.calls
234 .retain(|c| (c.at.start, &c.name) != (call.at.start, &call.name));
235 held.calls.push(call);
236 }
237 for one in chosen {
238 held.chosen.retain(|c| c.at.start != one.at.start);
239 held.chosen.push(one);
240 }
241 let old = self.draft.arguments.insert(node.to_string(), held);
242 self.draft.journal.push(Entry::Noted(node.to_string(), old));
243 }
244
245 pub fn arguments(&self, node: &str) -> Option<&Arguments> {
246 match self.draft.arguments.get(node) {
247 Some(held) => Some(held),
248 None if self.draft.hidden.contains(node) => None,
249 None => self.arguments.get(node),
250 }
251 }
252
253 pub fn ty(&self, n: NodeId) -> Ty {
254 self.at(n).ty
255 }
256
257 pub fn var(&self, n: NodeId) -> Var {
258 self.at(n).var
259 }
260
261 pub fn value(&self, n: NodeId) -> &Value {
262 &self.at(n).value
263 }
264
265 pub fn name(&self, n: NodeId) -> &str {
266 &self.at(n).name
267 }
268
269 pub fn grid(&self, n: NodeId) -> Grid {
270 self.at(n).grid
271 }
272
273 pub(crate) fn copy(&self, path: &str, grid: Grid) -> Option<NodeId> {
274 let key = (path.to_string(), grid);
275 match self.draft.copies.get(&key) {
276 Some(id) => Some(*id),
277 None if self.draft.hidden.contains(path) => None,
278 None => self
279 .copies
280 .get(path)
281 .and_then(|held| held.get(&grid))
282 .copied(),
283 }
284 }
285
286 pub(crate) fn copied(&mut self, path: &str, grid: Grid, id: NodeId) {
287 let key = (path.to_string(), grid);
288 let old = self.draft.copies.insert(key.clone(), id);
289 self.draft.journal.push(Entry::Copy(key, old));
290 }
291
292 pub(crate) fn opened(&mut self, path: &str) -> bool {
294 self.lowering.insert(path.to_string())
295 }
296
297 pub(crate) fn closed(&mut self, path: &str) {
298 self.lowering.remove(path);
299 }
300
301 pub fn at(&self, n: NodeId) -> &Node {
302 self.nodes[n.0 as usize]
303 .as_ref()
304 .expect("a node the typing holds")
305 }
306
307 pub fn ids(&self) -> impl Iterator<Item = NodeId> + '_ {
308 let held = self.nodes.iter().enumerate();
309 held.filter(|(_, n)| n.is_some())
310 .map(|(at, _)| NodeId(at as u32))
311 }
312
313 pub(crate) fn len(&self) -> usize {
314 self.nodes.len() - self.free.len()
315 }
316
317 pub fn id(&self, path: &str) -> Option<NodeId> {
318 match self.draft.by_path.get(path) {
319 Some(id) => Some(*id),
320 None if self.draft.hidden.contains(path) => None,
321 None => self.by_path.get(path).copied(),
322 }
323 }
324
325 pub fn paths(&self) -> impl Iterator<Item = (&str, NodeId)> {
327 let drafted = self.draft.by_path.iter();
328 let kept = self.by_path.iter().filter(|(path, _)| {
329 !self.draft.hidden.contains(*path) && !self.draft.by_path.contains_key(*path)
330 });
331 let all: BTreeMap<&String, &NodeId> = drafted.chain(kept).collect();
332 all.into_iter().map(|(p, id)| (p.as_str(), *id))
333 }
334
335 pub fn resolve(&self, path: &str) -> Result<NodeId, EngineError> {
337 if let Some(id) = self.id(path) {
338 return Ok(id);
339 }
340 let held = self.files.get(path).cloned().unwrap_or_default();
341 crate::instantiate::sole(
342 path,
343 held.into_iter()
344 .map(|id| (self.name(id).to_string(), id))
345 .collect(),
346 )
347 }
348
349 pub fn locate(&self, origin: Origin) -> Located {
351 let held = self.origins.get(origin.token() as usize).copied().flatten();
352 held.and_then(|(site, span)| {
353 let (name, _) = self.sites[site as usize].as_ref()?;
354 Some(Located::at(name.as_str(), span))
355 })
356 .unwrap_or_default()
357 }
358
359 pub(crate) fn mark(&mut self, node: &str, span: Option<sva_ast::ByteSpan>) -> Origin {
360 let site = match self.site_ids.get(node) {
361 Some(&site) => site,
362 None => {
363 let held = Some((node.to_string(), 0));
364 let site = match self.free_sites.pop() {
365 Some(site) => {
366 self.sites[site as usize] = held;
367 site
368 }
369 None => {
370 self.sites.push(held);
371 (self.sites.len() - 1) as u32
372 }
373 };
374 self.site_ids.insert(node.to_string(), site);
375 site
376 }
377 };
378 self.sites[site as usize].as_mut().expect("a site").1 += 1;
379 let token = match self.free_origins.pop() {
380 Some(token) => {
381 self.origins[token as usize] = Some((site, span));
382 token
383 }
384 None => {
385 self.origins.push(Some((site, span)));
386 (self.origins.len() - 1) as u32
387 }
388 };
389 let unit = self.unit();
390 self.draft.journal.push(Entry::Origin(token, unit));
391 Origin::new(token)
392 }
393
394 pub(crate) fn seed(&mut self, path: &str, held: Held, grid: Grid) -> NodeId {
395 if let Some(id) = self.id(path) {
396 return id;
397 }
398 self.begin(path, grid);
399 let id = self.push(
400 Node {
401 name: path.to_string(),
402 ty: Ty {
403 dual: held.is_closed_form(),
404 ..Ty::discrete(held, Codomain::Real)
405 },
406 var: Var::T,
407 value: Value::Op {
408 name: "loop".to_string(),
409 args: Vec::new(),
410 },
411 grid,
412 },
413 Some(path),
414 );
415 self.end();
416 self.pending.insert(id);
417 self.draft.journal.push(Entry::Pending(id));
418 id
419 }
420
421 pub(crate) fn pending(&self, id: NodeId) -> bool {
422 self.pending.contains(&id)
423 }
424
425 pub(crate) fn settle(&mut self, id: NodeId, node: Node) {
427 self.nodes[id.0 as usize] = Some(node);
428 self.pending.remove(&id);
429 self.numbers.0.get_mut().clear();
430 }
431
432 pub(crate) fn folded_number(&self, id: NodeId) -> Option<Option<C64>> {
433 self.numbers.0.borrow().get(&id).copied()
434 }
435
436 pub(crate) fn fold_number(&self, id: NodeId, number: Option<C64>) {
437 self.numbers.0.borrow_mut().insert(id, number);
438 }
439
440 pub(crate) fn alias(&mut self, path: &str, id: NodeId) {
441 let old = self.draft.by_path.insert(path.to_string(), id);
442 self.draft.journal.push(Entry::Path(path.to_string(), old));
443 }
444
445 pub(crate) fn push(&mut self, node: Node, path: Option<&str>) -> NodeId {
446 let id = match self.free.pop() {
447 Some(at) => {
448 self.nodes[at as usize] = Some(node);
449 NodeId(at)
450 }
451 None => {
452 self.nodes.push(Some(node));
453 NodeId((self.nodes.len() - 1) as u32)
454 }
455 };
456 let unit = self.unit();
457 self.draft.journal.push(Entry::Node(id, unit));
458 if let Some(path) = path {
459 self.alias(path, id);
460 }
461 id
462 }
463
464 pub(crate) fn begin(&mut self, path: &str, grid: Grid) {
466 self.making.push((path.to_string(), grid));
467 }
468
469 pub(crate) fn end(&mut self) {
470 self.making.pop();
471 }
472
473 fn unit(&self) -> (String, Grid) {
474 self.making
475 .last()
476 .cloned()
477 .unwrap_or_else(|| (String::new(), Grid::of(1)))
478 }
479
480 pub(crate) fn infer_closed_form(&self, form: &ClosedForm) -> Result<Ty, EngineError> {
481 infer(form, &Table(&self.nodes)).map_err(|r| {
482 EngineError::of_closed_form(
483 &r,
484 self.locate(r.origin),
485 "write the subterm inside sample(...) to leave A deliberately",
486 )
487 })
488 }
489}
490
491struct Table<'a>(&'a [Option<Node>]);
492
493impl Env for Table<'_> {
494 fn node(&self, id: NodeId) -> Ty {
495 self.0[id.0 as usize]
496 .as_ref()
497 .expect("a node the typing holds")
498 .ty
499 }
500
501 fn param(&self, _: ParamId) -> Ty {
503 Ty::form(Var::T, true, Codomain::Real)
504 }
505}
506
507pub fn infer_all(inst: &Instances, order: &Order) -> Result<Typing, EngineError> {
509 infer_over(inst, order, &BTreeMap::new())
510}
511
512pub(crate) fn infer_over(
514 inst: &Instances,
515 order: &Order,
516 stored: &BTreeMap<String, Arc<Stored>>,
517) -> Result<Typing, EngineError> {
518 let mut typing = Typing::default();
519 typing.lower(inst, &order.groups, stored)?;
520 typing.commit(inst);
521 Ok(typing)
522}
523
524impl Typing {
525 pub(crate) fn lower(
528 &mut self,
529 inst: &Instances,
530 groups: &[Vec<String>],
531 stored: &BTreeMap<String, Arc<Stored>>,
532 ) -> Result<(), EngineError> {
533 self.lowered.clear();
534 self.hide(groups.iter().flatten().cloned());
535 for group in groups {
536 match (crate::schedule::is_loop(inst, group), group.as_slice()) {
537 (false, [path]) if let Some(held) = stored.get(path) => {
538 self.begin(path, inst.grid());
539 self.push(standing(path, held), Some(path));
540 self.end();
541 }
542 (true, _) => settle_loop(self, inst, group)?,
543 (false, _) => {
544 for path in group {
545 lower::node(path, inst, self)?;
546 }
547 }
548 }
549 }
550 Ok(())
551 }
552}
553
554fn standing(path: &str, held: &Arc<Stored>) -> Node {
555 Node {
556 name: path.to_string(),
557 ty: Ty {
558 width: held.width,
559 rate: held.rate,
560 ..Ty::discrete(Held::Sampled, held.codomain)
561 },
562 var: Var::T,
563 value: Value::Stored(Arc::clone(held)),
564 grid: held.grid,
565 }
566}
567
568const SEEDS: [Held; 2] = [Held::Form(Var::T), Held::Sampled];
569
570fn settle_loop(typing: &mut Typing, inst: &Instances, group: &[String]) -> Result<(), EngineError> {
573 let mut seed = SEEDS[0];
574 let mut refusal = None;
575 for _ in 0..=SEEDS.len() {
576 let mark = typing.checkpoint();
577 for path in group {
578 typing.seed(path, seed, inst.grid());
579 }
580 let walked = group
581 .iter()
582 .try_for_each(|path| lower::node(path, inst, typing).map(|_| ()));
583 match walked {
584 Ok(()) => match held_by(typing, group) {
585 reached if reached == seed => return Ok(()),
586 reached => seed = reached,
587 },
588 Err(e) => {
589 refusal = Some(e);
590 seed = elsewhere(seed);
591 }
592 }
593 typing.rollback(mark);
594 }
595 Err(refusal.unwrap_or_else(|| mixed(group)))
596}
597
598fn elsewhere(seed: Held) -> Held {
600 match seed {
601 Held::Sampled => Held::Form(Var::T),
602 _ => Held::Sampled,
603 }
604}
605
606fn held_by(typing: &Typing, group: &[String]) -> Held {
608 let sampled = group
609 .iter()
610 .filter_map(|path| typing.id(path))
611 .any(|id| !typing.ty(id).is_closed_form());
612 match sampled {
613 true => Held::Sampled,
614 false => Held::Form(Var::T),
615 }
616}
617
618fn mixed(group: &[String]) -> EngineError {
620 let at = group.first().map_or("", String::as_str);
621 EngineError::refused(Diagnostic {
622 code: "type.samples_in_closed_form".to_string(),
623 message: format!(
624 "the loop over `{}` holds a closed form and samples at once.",
625 group.join("`, `")
626 ),
627 location: Located::at(at, None),
628 help: "write sample(...) on the members that are closed forms, so the whole loop runs \
629 on the grid"
630 .to_string(),
631 })
632}