1use std::cell::RefCell;
4use std::collections::{BTreeMap, BTreeSet};
5use std::sync::Arc;
6
7use sva_formula::filter::Shape;
8use sva_formula::{C64, ClosedForm, Codomain, Env, Held, NodeId, Origin, ParamId, Ty, Var, infer};
9use sva_samples::Params;
10
11use crate::arguments::{Arguments, Called, Chosen};
12use crate::cache::Stored;
13use crate::cast::Cast;
14use crate::error::{Diagnostic, EngineError, Located};
15use crate::instantiate::Instances;
16use crate::lower;
17use crate::schedule::Order;
18use crate::time::Grid;
19
20#[derive(Clone, Debug, PartialEq)]
22pub enum Value {
23 ClosedForm(ClosedForm),
24 Cast(Cast, NodeId),
25 Op {
26 name: String,
27 args: Vec<NodeId>,
28 },
29 SelfAt {
30 at: When,
31 },
32 Read {
33 source: NodeId,
34 at: When,
35 site: Origin,
36 },
37 Noise(u64),
38 Filter {
40 shape: Shape,
41 x: NodeId,
42 cutoff: NodeId,
43 q: NodeId,
44 gain: NodeId,
45 },
46 Solver {
47 params: Box<Params>,
48 varying: Vec<(&'static str, NodeId)>,
49 },
50 Stored(Arc<Stored>),
52}
53
54#[derive(Clone, Debug, PartialEq)]
57pub enum When {
58 At(crate::time::Affine),
59 Moving(NodeId),
60 Index(crate::index::Index),
61 Step(Step),
62}
63
64#[derive(Clone, Debug, PartialEq)]
67pub enum Step {
68 Index(crate::index::Index),
69 Nearest(NodeId, crate::index::Round),
70 Add(Vec<Step>),
71 Neg(Box<Step>),
72 Mul(Vec<Step>),
73}
74
75impl Step {
76 pub(crate) fn times(&self, out: &mut Vec<NodeId>) {
78 match self {
79 Step::Index(_) => {}
80 Step::Nearest(time, _) => out.push(*time),
81 Step::Add(parts) | Step::Mul(parts) => parts.iter().for_each(|p| p.times(out)),
82 Step::Neg(p) => p.times(out),
83 }
84 }
85}
86
87impl When {
88 pub(crate) fn moving(&self) -> Vec<NodeId> {
90 let mut out = Vec::new();
91 match self {
92 When::Moving(id) => out.push(*id),
93 When::Step(step) => step.times(&mut out),
94 When::At(_) | When::Index(_) => {}
95 }
96 out
97 }
98}
99
100#[derive(Clone, Debug, PartialEq)]
101pub struct Node {
102 pub name: String,
103 pub ty: Ty,
104 pub var: Var,
105 pub value: Value,
106 pub grid: Grid,
107}
108
109#[derive(Clone, Debug, Default, PartialEq)]
110pub struct Typing {
111 nodes: Vec<Node>,
112 arguments: BTreeMap<String, Arguments>,
113 by_path: BTreeMap<String, NodeId>,
114 copies: BTreeMap<(String, Grid), NodeId>,
115 lowering: BTreeSet<String>,
116 files: BTreeMap<String, Vec<NodeId>>,
117 origins: Vec<(u32, Option<sva_ast::ByteSpan>)>,
118 sites: Vec<String>,
119 site_ids: BTreeMap<String, u32>,
120 pending: BTreeSet<NodeId>,
121 indices: u32,
122 sum: Option<(NodeId, Vec<SumSlot>)>,
123 numbers: Numbers,
124 lowered: Vec<String>,
126}
127
128#[derive(Clone, Debug, Default)]
130struct Numbers(RefCell<BTreeMap<NodeId, Option<C64>>>);
131
132impl PartialEq for Numbers {
133 fn eq(&self, _: &Numbers) -> bool {
134 true
135 }
136}
137
138#[derive(Clone, Copy, Debug, PartialEq)]
140pub(crate) enum SumSlot {
141 Node(NodeId),
142 Retired(sva_samples::Extent),
143}
144
145impl Typing {
146 pub(crate) fn name_sum(&mut self, node: NodeId, slots: Vec<SumSlot>) {
148 self.sum = Some((node, slots));
149 }
150
151 pub(crate) fn retires(&self) -> bool {
153 self.sum
154 .as_ref()
155 .is_some_and(|(_, slots)| slots.iter().any(|slot| matches!(slot, SumSlot::Retired(_))))
156 }
157
158 pub(crate) fn sum_slots(&self, node: NodeId) -> Option<&[SumSlot]> {
159 self.sum
160 .as_ref()
161 .filter(|(held, _)| *held == node)
162 .map(|(_, slots)| slots.as_slice())
163 }
164
165 pub(crate) fn lowering(&mut self, path: &str) {
166 self.lowered.push(path.to_string());
167 }
168
169 pub(crate) fn lowered(&self) -> &[String] {
170 &self.lowered
171 }
172
173 pub(crate) fn next_index(&mut self) -> sva_formula::IndexId {
174 self.indices += 1;
175 sva_formula::IndexId(self.indices)
176 }
177
178 pub(crate) fn note(&mut self, node: &str, call: Option<Called>, chosen: Vec<Chosen>) {
180 let held = self
181 .arguments
182 .entry(node.to_string())
183 .or_insert_with(|| Arguments {
184 node: node.to_string(),
185 ..Arguments::default()
186 });
187 if let Some(call) = call {
188 held.calls
189 .retain(|c| (c.at.start, &c.name) != (call.at.start, &call.name));
190 held.calls.push(call);
191 }
192 for one in chosen {
193 held.chosen.retain(|c| c.at.start != one.at.start);
194 held.chosen.push(one);
195 }
196 }
197
198 pub fn arguments(&self, node: &str) -> Option<&Arguments> {
199 self.arguments.get(node)
200 }
201
202 pub fn ty(&self, n: NodeId) -> Ty {
203 self.at(n).ty
204 }
205
206 pub fn var(&self, n: NodeId) -> Var {
207 self.at(n).var
208 }
209
210 pub fn value(&self, n: NodeId) -> &Value {
211 &self.at(n).value
212 }
213
214 pub fn name(&self, n: NodeId) -> &str {
215 &self.at(n).name
216 }
217
218 pub fn grid(&self, n: NodeId) -> Grid {
219 self.at(n).grid
220 }
221
222 pub(crate) fn copy(&self, path: &str, grid: Grid) -> Option<NodeId> {
223 self.copies.get(&(path.to_string(), grid)).copied()
224 }
225
226 pub(crate) fn copied(&mut self, path: &str, grid: Grid, id: NodeId) {
227 self.copies.insert((path.to_string(), grid), id);
228 }
229
230 pub(crate) fn opened(&mut self, path: &str) -> bool {
232 self.lowering.insert(path.to_string())
233 }
234
235 pub(crate) fn closed(&mut self, path: &str) {
236 self.lowering.remove(path);
237 }
238
239 pub fn at(&self, n: NodeId) -> &Node {
240 &self.nodes[n.0 as usize]
241 }
242
243 pub(crate) fn len(&self) -> usize {
244 self.nodes.len()
245 }
246
247 pub fn id(&self, path: &str) -> Option<NodeId> {
248 self.by_path.get(path).copied()
249 }
250
251 pub fn paths(&self) -> impl Iterator<Item = (&str, NodeId)> {
252 self.by_path.iter().map(|(p, id)| (p.as_str(), *id))
253 }
254
255 pub fn resolve(&self, path: &str) -> Result<NodeId, EngineError> {
257 if let Some(id) = self.id(path) {
258 return Ok(id);
259 }
260 let held = self.files.get(path).cloned().unwrap_or_default();
261 crate::instantiate::sole(
262 path,
263 held.into_iter()
264 .map(|id| (self.name(id).to_string(), id))
265 .collect(),
266 )
267 }
268
269 pub fn locate(&self, origin: Origin) -> Located {
271 self.origins
272 .get(origin.token() as usize)
273 .map(|&(site, span)| Located::at(self.sites[site as usize].as_str(), span))
274 .unwrap_or_default()
275 }
276
277 pub(crate) fn mark(&mut self, node: &str, span: Option<sva_ast::ByteSpan>) -> Origin {
278 let site = match self.site_ids.get(node) {
279 Some(&site) => site,
280 None => {
281 let site = self.sites.len() as u32;
282 self.sites.push(node.to_string());
283 self.site_ids.insert(node.to_string(), site);
284 site
285 }
286 };
287 self.origins.push((site, span));
288 Origin::new((self.origins.len() - 1) as u32)
289 }
290
291 pub(crate) fn seed(&mut self, path: &str, held: Held, grid: Grid) -> NodeId {
292 if let Some(id) = self.id(path) {
293 return id;
294 }
295 let id = self.push(
296 Node {
297 name: path.to_string(),
298 ty: Ty {
299 dual: held.is_closed_form(),
300 ..Ty::discrete(held, Codomain::Real)
301 },
302 var: Var::T,
303 value: Value::Op {
304 name: "loop".to_string(),
305 args: Vec::new(),
306 },
307 grid,
308 },
309 Some(path),
310 );
311 self.pending.insert(id);
312 id
313 }
314
315 pub(crate) fn pending(&self, id: NodeId) -> bool {
316 self.pending.contains(&id)
317 }
318
319 pub(crate) fn settle(&mut self, id: NodeId, node: Node) {
320 self.nodes[id.0 as usize] = node;
321 self.pending.remove(&id);
322 self.numbers.0.get_mut().clear();
323 }
324
325 pub(crate) fn folded_number(&self, id: NodeId) -> Option<Option<C64>> {
326 self.numbers.0.borrow().get(&id).copied()
327 }
328
329 pub(crate) fn fold_number(&self, id: NodeId, number: Option<C64>) {
330 self.numbers.0.borrow_mut().insert(id, number);
331 }
332
333 pub(crate) fn alias(&mut self, path: &str, id: NodeId) {
334 self.by_path.insert(path.to_string(), id);
335 }
336
337 pub(crate) fn push(&mut self, node: Node, path: Option<&str>) -> NodeId {
338 let id = NodeId(self.nodes.len() as u32);
339 self.nodes.push(node);
340 if let Some(path) = path {
341 self.by_path.insert(path.to_string(), id);
342 }
343 id
344 }
345
346 pub(crate) fn infer_closed_form(&self, form: &ClosedForm) -> Result<Ty, EngineError> {
347 infer(form, &Table(&self.nodes)).map_err(|r| {
348 EngineError::of_closed_form(
349 &r,
350 self.locate(r.origin),
351 "write the subterm inside sample(...) to leave A deliberately",
352 )
353 })
354 }
355}
356
357struct Table<'a>(&'a [Node]);
358
359impl Env for Table<'_> {
360 fn node(&self, id: NodeId) -> Ty {
361 self.0[id.0 as usize].ty
362 }
363
364 fn param(&self, _: ParamId) -> Ty {
366 Ty::form(Var::T, true, Codomain::Real)
367 }
368}
369
370pub fn infer_all(inst: &Instances, order: &Order) -> Result<Typing, EngineError> {
372 infer_over(inst, order, &BTreeMap::new())
373}
374
375pub(crate) fn infer_over(
377 inst: &Instances,
378 order: &Order,
379 stored: &BTreeMap<String, Arc<Stored>>,
380) -> Result<Typing, EngineError> {
381 let mut typing = Typing::default();
382 for group in &order.groups {
383 match (order.is_loop(group), group.as_slice()) {
384 (false, [path]) if let Some(held) = stored.get(path) => {
385 typing.push(standing(path, held), Some(path));
386 }
387 (true, _) => settle_loop(&mut typing, inst, group)?,
388 (false, _) => {
389 for path in group {
390 lower::node(path, inst, &mut typing)?;
391 }
392 }
393 }
394 }
395 name_files(&mut typing, inst);
396 Ok(typing)
397}
398
399fn standing(path: &str, held: &Arc<Stored>) -> Node {
400 Node {
401 name: path.to_string(),
402 ty: Ty {
403 width: held.width,
404 rate: held.rate,
405 ..Ty::discrete(Held::Sampled, held.codomain)
406 },
407 var: Var::T,
408 value: Value::Stored(Arc::clone(held)),
409 grid: held.grid,
410 }
411}
412
413const SEEDS: [Held; 2] = [Held::Form(Var::T), Held::Sampled];
414
415fn settle_loop(typing: &mut Typing, inst: &Instances, group: &[String]) -> Result<(), EngineError> {
418 let mut seed = SEEDS[0];
419 let mut refusal = None;
420 for _ in 0..=SEEDS.len() {
421 let mut attempt = typing.clone();
422 for path in group {
423 attempt.seed(path, seed, inst.grid());
424 }
425 let walked = group
426 .iter()
427 .try_for_each(|path| lower::node(path, inst, &mut attempt).map(|_| ()));
428 match walked {
429 Ok(()) => match held_by(&attempt, group) {
430 reached if reached == seed => {
431 *typing = attempt;
432 return Ok(());
433 }
434 reached => seed = reached,
435 },
436 Err(e) => {
437 refusal = Some(e);
438 seed = elsewhere(seed);
439 }
440 }
441 }
442 Err(refusal.unwrap_or_else(|| mixed(group)))
443}
444
445fn elsewhere(seed: Held) -> Held {
447 match seed {
448 Held::Sampled => Held::Form(Var::T),
449 _ => Held::Sampled,
450 }
451}
452
453fn held_by(typing: &Typing, group: &[String]) -> Held {
455 let sampled = group
456 .iter()
457 .filter_map(|path| typing.id(path))
458 .any(|id| !typing.ty(id).is_closed_form());
459 match sampled {
460 true => Held::Sampled,
461 false => Held::Form(Var::T),
462 }
463}
464
465fn mixed(group: &[String]) -> EngineError {
467 let at = group.first().map_or("", String::as_str);
468 EngineError::refused(Diagnostic {
469 code: "type.samples_in_closed_form".to_string(),
470 message: format!(
471 "the loop over `{}` holds a closed form and samples at once.",
472 group.join("`, `")
473 ),
474 location: Located::at(at, None),
475 help: "write sample(...) on the members that are closed forms, so the whole loop runs \
476 on the grid"
477 .to_string(),
478 })
479}
480
481fn name_files(typing: &mut Typing, inst: &Instances) {
483 let files: Vec<String> = inst
484 .paths()
485 .filter_map(|p| inst.origin(p))
486 .map(str::to_string)
487 .collect();
488 for file in files {
489 if typing.files.contains_key(&file) {
490 continue;
491 }
492 let held: Vec<NodeId> = inst
493 .instances_of(&file)
494 .filter_map(|p| typing.id(&p))
495 .collect();
496 if let ([only], None) = (held.as_slice(), typing.id(&file)) {
497 typing.alias(&file, *only);
498 }
499 typing.files.insert(file, held);
500 }
501}