1use std::collections::{BTreeMap, BTreeSet, HashMap};
4
5use sva_ast::{Arg, Expr};
6use sva_formula::NodeId;
7
8use crate::cast::Cast;
9use crate::error::EngineError;
10use crate::instantiate::{Cx, Instances, Node};
11use crate::query::Ask;
12use crate::typing::{Typing, Value};
13
14fn direct_refs(inst: &Instances, e: &Expr, cx: Cx, out: &mut Vec<String>) {
16 if inst
17 .follow(e, cx, |e2, cx2| direct_refs(inst, e2, cx2, out))
18 .is_some()
19 {
20 return;
21 }
22 match inst.node(e, cx) {
23 Node::Lit(_) | Node::Name(_) => {}
24 Node::Bin(_, l, r) => {
25 direct_refs(inst, l, cx, out);
26 direct_refs(inst, r, cx, out);
27 }
28 Node::Call { args, .. } => {
29 for a in args {
30 let (Arg::Pos(x) | Arg::Named(_, x)) = a;
31 direct_refs(inst, x, cx, out);
32 }
33 }
34 Node::Read { path, arg, .. } => {
35 out.push(path.to_string());
36 direct_refs(inst, arg, cx, out);
37 }
38 Node::Own { arg, .. } => direct_refs(inst, arg, cx, out),
39 }
40}
41
42pub struct Order {
44 pub groups: Vec<Vec<String>>,
45 deps: BTreeMap<String, Vec<String>>,
46}
47
48impl Order {
49 pub fn deps(&self, path: &str) -> &[String] {
50 self.deps.get(path).map_or(&[], Vec::as_slice)
51 }
52
53 pub fn is_loop(&self, group: &[String]) -> bool {
55 match group {
56 [only] => self.deps(only).iter().any(|d| d == only),
57 _ => true,
58 }
59 }
60}
61
62pub fn direct_deps(inst: &Instances, path: &str) -> Result<Vec<String>, EngineError> {
63 let (e, cx) = inst
64 .at(path)
65 .ok_or_else(|| EngineError::UnknownNode(path.to_string()))?;
66 let mut out = Vec::new();
67 direct_refs(inst, e, cx, &mut out);
68 out.sort();
69 out.dedup();
70 for target in &out {
71 if !inst.holds(target) {
72 return Err(EngineError::UnknownNode(target.clone()));
73 }
74 }
75 Ok(out)
76}
77
78struct Frame {
79 node: String,
80 refs: Vec<String>,
81 idx: usize,
82}
83
84pub fn schedule_from(inst: &Instances, roots: &[String]) -> Result<Order, EngineError> {
86 let mut walk = Walk {
87 inst,
88 deps: BTreeMap::new(),
89 index: HashMap::new(),
90 low: HashMap::new(),
91 open: Vec::new(),
92 next: 0,
93 groups: Vec::new(),
94 };
95 for root in roots {
96 walk.from(root)?;
97 }
98 Ok(Order {
99 groups: walk.groups,
100 deps: walk.deps,
101 })
102}
103
104struct Walk<'a> {
105 inst: &'a Instances<'a>,
106 deps: BTreeMap<String, Vec<String>>,
107 index: HashMap<String, usize>,
108 low: HashMap<String, usize>,
109 open: Vec<String>,
110 next: usize,
111 groups: Vec<Vec<String>>,
112}
113
114impl Walk<'_> {
115 fn from(&mut self, root: &str) -> Result<(), EngineError> {
116 if !self.inst.holds(root) {
117 return Err(EngineError::UnknownNode(root.to_string()));
118 }
119 if self.index.contains_key(root) {
120 return Ok(());
121 }
122 self.index.insert(root.to_string(), self.next);
123 self.low.insert(root.to_string(), self.next);
124 self.next += 1;
125 self.open.push(root.to_string());
126 let seed = direct_deps(self.inst, root)?;
127 self.deps.insert(root.to_string(), seed.clone());
128 let mut stack = vec![Frame {
129 node: root.to_string(),
130 refs: seed,
131 idx: 0,
132 }];
133
134 while let Some(frame) = stack.last_mut() {
135 if frame.idx < frame.refs.len() {
136 let target = frame.refs[frame.idx].clone();
137 let node = frame.node.clone();
138 frame.idx += 1;
139 match self.index.get(&target).copied() {
140 None => {
141 self.index.insert(target.clone(), self.next);
142 self.low.insert(target.clone(), self.next);
143 self.next += 1;
144 self.open.push(target.clone());
145 let refs = direct_deps(self.inst, &target)?;
146 self.deps.insert(target.clone(), refs.clone());
147 stack.push(Frame {
148 node: target,
149 refs,
150 idx: 0,
151 });
152 }
153 Some(at) if self.open.contains(&target) => {
154 let mine = self.low[&node];
155 self.low.insert(node, mine.min(at));
156 }
157 Some(_) => {}
158 }
159 continue;
160 }
161
162 let node = frame.node.clone();
163 let mine = self.low[&node];
164 stack.pop();
165 if let Some(parent) = stack.last() {
166 let above = self.low[&parent.node];
167 self.low.insert(parent.node.clone(), above.min(mine));
168 }
169 if mine == self.index[&node] {
170 let at = self
171 .open
172 .iter()
173 .rposition(|n| *n == node)
174 .expect("a root of its group is still open");
175 let mut group = self.open.split_off(at);
176 group.sort();
177 self.groups.push(group);
178 }
179 }
180 Ok(())
181 }
182}
183
184#[derive(Clone, Debug, Default, PartialEq)]
186pub struct Schedule {
187 pub materialize: Vec<NodeId>,
188 pub wanted: Vec<NodeId>,
190 pub symbolic: Vec<NodeId>,
191 pub compose: Vec<NodeId>,
194}
195
196pub fn plan(typing: &Typing, order: &Order, root: NodeId, asks: &[Ask]) -> Schedule {
198 let mut wanted: BTreeSet<NodeId> = BTreeSet::new();
199 let mut compose: Vec<NodeId> = Vec::new();
200 let audio = asks.is_empty();
201 for ask in asks {
202 let Some(id) = typing.id(&ask.node) else {
203 continue;
204 };
205 if matches!(
207 ask.representation,
208 crate::query::Representation::Bindings
209 | crate::query::Representation::Arguments
210 | crate::query::Representation::Flops
211 ) {
212 continue;
213 }
214 match ask.representation.consumes(typing.ty(id).is_closed_form()) {
215 sva_samples::Consumes::ClosedForm if !compose.contains(&id) => compose.push(id),
216 sva_samples::Consumes::ClosedForm => {}
217 _ => {
218 wanted.insert(id);
219 }
220 }
221 if let crate::query::Representation::Ledger { depth } = ask.representation {
222 attributed(typing, id, depth, &mut wanted);
223 }
224 }
225 if audio || !wanted.is_empty() {
226 wanted.insert(root);
227 }
228
229 let mut reached: BTreeSet<NodeId> = BTreeSet::new();
230 let mut work: Vec<NodeId> = wanted.iter().copied().collect();
231 while let Some(id) = work.pop() {
232 if !reached.insert(id) {
233 continue;
234 }
235 work.extend(materialized_operands(typing, id));
236 }
237
238 let held: Vec<NodeId> = order
239 .groups
240 .concat()
241 .iter()
242 .filter_map(|path| typing.id(path))
243 .collect();
244 let mut materialize: Vec<NodeId> = Vec::new();
245 let mut seen: BTreeSet<NodeId> = BTreeSet::new();
246 for id in held {
247 for member in dependencies_first(typing, id, &mut seen) {
248 if reached.contains(&member) {
249 materialize.push(member);
250 }
251 }
252 }
253 let symbolic = typing
254 .paths()
255 .map(|(_, id)| id)
256 .filter(|id| !materialize.contains(id))
257 .collect();
258 compose.retain(|id| !materialize.contains(id));
259 Schedule {
260 wanted: materialize
261 .iter()
262 .copied()
263 .filter(|id| wanted.contains(id))
264 .collect(),
265 materialize,
266 symbolic,
267 compose,
268 }
269}
270
271fn attributed(typing: &Typing, id: NodeId, depth: usize, wanted: &mut BTreeSet<NodeId>) {
273 let mut seen = BTreeSet::from([id]);
275 let mut level = vec![id];
276 for _ in 0..depth {
277 let mut next = Vec::new();
278 for held in level {
279 for operand in read_operands(typing, held) {
280 if seen.insert(operand) {
281 wanted.insert(operand);
282 next.push(operand);
283 }
284 }
285 }
286 level = next;
287 }
288}
289
290pub(crate) fn holds_self(typing: &Typing, id: NodeId, seen: &mut BTreeSet<NodeId>) -> bool {
293 if !seen.insert(id) {
294 return false;
295 }
296 match typing.value(id) {
297 Value::SelfAt(_) => true,
298 Value::Cast(Cast::Sample, _) | Value::Read { .. } => false,
299 Value::Cast(_, source) => holds_self(typing, *source, seen),
300 Value::Op { args, .. } => args.iter().any(|a| holds_self(typing, *a, seen)),
301 Value::Filter {
302 x, cutoff, q, gain, ..
303 } => [x, cutoff, q, gain]
304 .into_iter()
305 .any(|operand| holds_self(typing, *operand, seen)),
306 Value::ClosedForm(_) | Value::Solver(_) | Value::Grid(_) => false,
307 }
308}
309
310pub(crate) fn materialized_operands(typing: &Typing, id: NodeId) -> Vec<NodeId> {
312 let sampled = |set: Vec<NodeId>| -> Vec<NodeId> {
313 let mut out = Vec::new();
314 for op in set {
315 if typing.ty(op).is_closed_form() || matches!(typing.value(op), Value::Grid(_)) {
316 continue;
317 }
318 match holds_self(typing, op, &mut BTreeSet::new()) {
319 true => out.extend(materialized_operands(typing, op)),
320 false => out.push(op),
321 }
322 }
323 out
324 };
325 match typing.value(id) {
326 Value::ClosedForm(_) | Value::SelfAt(_) | Value::Solver(_) | Value::Grid(_) => Vec::new(),
327 Value::Cast(Cast::Sample, source) => vec![*source],
328 Value::Read { source, .. } => vec![*source],
329 Value::Cast(_, source) => sampled(vec![*source]),
330 Value::Op { args, .. } => sampled(args.clone()),
331 Value::Filter {
332 x, cutoff, q, gain, ..
333 } => sampled(vec![*x, *cutoff, *q, *gain]),
334 }
335}
336
337pub(crate) fn forks(typing: &Typing, held: &[NodeId]) -> BTreeSet<NodeId> {
340 let mut readers: BTreeMap<NodeId, usize> = BTreeMap::new();
341 for id in held {
342 let mut read = materialized_operands(typing, *id);
343 read.sort_unstable();
344 read.dedup();
345 for operand in read {
346 *readers.entry(operand).or_default() += 1;
347 }
348 }
349 readers
350 .into_iter()
351 .filter(|(id, count)| *count >= 2 && held.contains(id))
352 .map(|(id, _)| id)
353 .collect()
354}
355
356pub(crate) fn dependencies_first(
358 typing: &Typing,
359 id: NodeId,
360 seen: &mut BTreeSet<NodeId>,
361) -> Vec<NodeId> {
362 if !seen.insert(id) {
363 return Vec::new();
364 }
365 let mut out = Vec::new();
366 for operand in materialized_operands(typing, id) {
367 out.extend(dependencies_first(typing, operand, seen));
368 }
369 out.push(id);
370 out
371}
372
373pub(crate) fn read_operands(typing: &Typing, id: NodeId) -> Vec<NodeId> {
375 match typing.value(id) {
376 Value::ClosedForm(form) => crate::refs::nodes_in(&form.body),
377 _ if typing.ty(id).is_closed_form() => {
378 let mut out = Vec::new();
379 reads_under(typing, id, &mut BTreeSet::new(), &mut out);
380 out.dedup();
381 out
382 }
383 _ => materialized_operands(typing, id),
384 }
385}
386
387fn reads_under(typing: &Typing, id: NodeId, seen: &mut BTreeSet<NodeId>, out: &mut Vec<NodeId>) {
389 if !seen.insert(id) {
390 return;
391 }
392 match typing.value(id) {
393 Value::ClosedForm(form) => out.extend(crate::refs::nodes_in(&form.body)),
394 Value::Read { source, .. } => out.push(*source),
395 Value::Cast(_, source) => read_through(typing, *source, seen, out),
396 Value::Op { args, .. } => {
397 for arg in args {
398 read_through(typing, *arg, seen, out);
399 }
400 }
401 Value::Filter {
402 x, cutoff, q, gain, ..
403 } => {
404 for operand in [x, cutoff, q, gain] {
405 read_through(typing, *operand, seen, out);
406 }
407 }
408 Value::SelfAt(_) | Value::Grid(_) | Value::Solver(_) => {}
409 }
410}
411
412fn read_through(typing: &Typing, id: NodeId, seen: &mut BTreeSet<NodeId>, out: &mut Vec<NodeId>) {
414 let Value::ClosedForm(_) = typing.value(id) else {
415 return reads_under(typing, id, seen, out);
416 };
417 if seen.insert(id) {
418 out.push(id);
419 }
420}