1use std::collections::{BTreeMap, BTreeSet};
4
5use sva_formula::filter::Shape;
6use sva_formula::{ClosedForm, Codomain, Env, Held, NodeId, Origin, ParamId, Ty, Var, infer};
7use sva_samples::Params;
8
9use crate::cast::Cast;
10use crate::error::{Diagnostic, EngineError, Located};
11use crate::instantiate::Instances;
12use crate::loops::Delay;
13use crate::lower;
14use crate::offset::Offset;
15use crate::schedule::Order;
16
17#[derive(Clone, Debug, PartialEq)]
19pub enum Value {
20 ClosedForm(ClosedForm),
21 Cast(Cast, NodeId),
22 Op {
23 name: String,
24 args: Vec<NodeId>,
25 },
26 SelfAt(Delay),
28 Grid(f64),
30 Read {
31 source: NodeId,
32 at: Offset,
33 site: Origin,
34 },
35 Filter {
37 shape: Shape,
38 x: NodeId,
39 cutoff: NodeId,
40 q: NodeId,
41 gain: NodeId,
42 },
43 Solver(Box<Params>),
44}
45
46#[derive(Clone, Debug, PartialEq)]
47pub struct Node {
48 pub name: String,
49 pub ty: Ty,
50 pub var: Var,
51 pub value: Value,
52}
53
54#[derive(Clone, Debug, Default, PartialEq)]
55pub struct Typing {
56 nodes: Vec<Node>,
57 by_path: BTreeMap<String, NodeId>,
58 files: BTreeMap<String, Vec<NodeId>>,
59 origins: Vec<Located>,
60 pending: BTreeSet<NodeId>,
61 indices: u32,
62}
63
64impl Typing {
65 pub(crate) fn next_index(&mut self) -> sva_formula::IndexId {
66 self.indices += 1;
67 sva_formula::IndexId(self.indices)
68 }
69
70 pub fn ty(&self, n: NodeId) -> Ty {
71 self.at(n).ty
72 }
73
74 pub fn var(&self, n: NodeId) -> Var {
75 self.at(n).var
76 }
77
78 pub fn value(&self, n: NodeId) -> &Value {
79 &self.at(n).value
80 }
81
82 pub fn name(&self, n: NodeId) -> &str {
83 &self.at(n).name
84 }
85
86 pub fn at(&self, n: NodeId) -> &Node {
87 &self.nodes[n.0 as usize]
88 }
89
90 pub fn id(&self, path: &str) -> Option<NodeId> {
91 self.by_path.get(path).copied()
92 }
93
94 pub fn paths(&self) -> impl Iterator<Item = (&str, NodeId)> {
95 self.by_path.iter().map(|(p, id)| (p.as_str(), *id))
96 }
97
98 pub fn resolve(&self, path: &str) -> Result<NodeId, EngineError> {
100 if let Some(id) = self.id(path) {
101 return Ok(id);
102 }
103 let held = self.files.get(path).cloned().unwrap_or_default();
104 crate::instantiate::sole(
105 path,
106 held.into_iter()
107 .map(|id| (self.name(id).to_string(), id))
108 .collect(),
109 )
110 }
111
112 pub fn locate(&self, origin: Origin) -> Located {
114 self.origins
115 .get(origin.token() as usize)
116 .cloned()
117 .unwrap_or_default()
118 }
119
120 pub(crate) fn mark(&mut self, at: Located) -> Origin {
121 self.origins.push(at);
122 Origin::new((self.origins.len() - 1) as u32)
123 }
124
125 pub(crate) fn seed(&mut self, path: &str, held: Held) -> NodeId {
126 if let Some(id) = self.id(path) {
127 return id;
128 }
129 let id = self.push(
130 Node {
131 name: path.to_string(),
132 ty: Ty {
133 dual: held.is_closed_form(),
134 ..Ty::discrete(held, Codomain::Real)
135 },
136 var: Var::T,
137 value: Value::Op {
138 name: "loop".to_string(),
139 args: Vec::new(),
140 },
141 },
142 Some(path),
143 );
144 self.pending.insert(id);
145 id
146 }
147
148 pub(crate) fn pending(&self, id: NodeId) -> bool {
149 self.pending.contains(&id)
150 }
151
152 pub(crate) fn settle(&mut self, id: NodeId, node: Node) {
153 self.nodes[id.0 as usize] = node;
154 self.pending.remove(&id);
155 }
156
157 pub(crate) fn alias(&mut self, path: &str, id: NodeId) {
158 self.by_path.insert(path.to_string(), id);
159 }
160
161 pub(crate) fn push(&mut self, node: Node, path: Option<&str>) -> NodeId {
162 let id = NodeId(self.nodes.len() as u32);
163 self.nodes.push(node);
164 if let Some(path) = path {
165 self.by_path.insert(path.to_string(), id);
166 }
167 id
168 }
169
170 pub(crate) fn infer_closed_form(&self, form: &ClosedForm) -> Result<Ty, EngineError> {
171 infer(form, &Table(&self.nodes)).map_err(|r| {
172 EngineError::of_closed_form(
173 &r,
174 self.locate(r.origin),
175 "write the subterm inside sample(...) to leave A deliberately",
176 )
177 })
178 }
179}
180
181struct Table<'a>(&'a [Node]);
182
183impl Env for Table<'_> {
184 fn node(&self, id: NodeId) -> Ty {
185 self.0[id.0 as usize].ty
186 }
187
188 fn param(&self, _: ParamId) -> Ty {
190 Ty::form(Var::T, true, Codomain::Real)
191 }
192}
193
194pub fn infer_all(inst: &Instances, order: &Order) -> Result<Typing, EngineError> {
196 let mut typing = Typing::default();
197 for group in &order.groups {
198 match order.is_loop(group) {
199 true => settle_loop(&mut typing, inst, group)?,
200 false => {
201 for path in group {
202 lower::node(path, inst, &mut typing)?;
203 }
204 }
205 }
206 }
207 name_files(&mut typing, inst);
208 Ok(typing)
209}
210
211const SEEDS: [Held; 2] = [Held::Form(Var::T), Held::Sampled];
212
213fn settle_loop(typing: &mut Typing, inst: &Instances, group: &[String]) -> Result<(), EngineError> {
216 let mut seed = SEEDS[0];
217 let mut refusal = None;
218 for _ in 0..=SEEDS.len() {
219 let mut attempt = typing.clone();
220 for path in group {
221 attempt.seed(path, seed);
222 }
223 let walked = group
224 .iter()
225 .try_for_each(|path| lower::node(path, inst, &mut attempt).map(|_| ()));
226 match walked {
227 Ok(()) => match held_by(&attempt, group) {
228 reached if reached == seed => {
229 *typing = attempt;
230 return Ok(());
231 }
232 reached => seed = reached,
233 },
234 Err(e) => {
235 refusal = Some(e);
236 seed = elsewhere(seed);
237 }
238 }
239 }
240 Err(refusal.unwrap_or_else(|| mixed(group)))
241}
242
243fn elsewhere(seed: Held) -> Held {
245 match seed {
246 Held::Sampled => Held::Form(Var::T),
247 _ => Held::Sampled,
248 }
249}
250
251fn held_by(typing: &Typing, group: &[String]) -> Held {
253 let sampled = group
254 .iter()
255 .filter_map(|path| typing.id(path))
256 .any(|id| !typing.ty(id).is_closed_form());
257 match sampled {
258 true => Held::Sampled,
259 false => Held::Form(Var::T),
260 }
261}
262
263fn mixed(group: &[String]) -> EngineError {
265 let at = group.first().map_or("", String::as_str);
266 EngineError::refused(Diagnostic {
267 code: "type.samples_in_closed_form".to_string(),
268 message: format!(
269 "the loop over `{}` holds a closed form and samples at once.",
270 group.join("`, `")
271 ),
272 location: Located::at(at, None),
273 help: "write sample(...) on the members that are closed forms, so the whole loop runs \
274 on the grid"
275 .to_string(),
276 })
277}
278
279fn name_files(typing: &mut Typing, inst: &Instances) {
281 let files: Vec<String> = inst
282 .paths()
283 .filter_map(|p| inst.origin(p))
284 .map(str::to_string)
285 .collect();
286 for file in files {
287 if typing.files.contains_key(&file) {
288 continue;
289 }
290 let held: Vec<NodeId> = inst
291 .instances_of(&file)
292 .filter_map(|p| typing.id(&p))
293 .collect();
294 if let ([only], None) = (held.as_slice(), typing.id(&file)) {
295 typing.alias(&file, *only);
296 }
297 typing.files.insert(file, held);
298 }
299}