1use std::collections::{BTreeSet, HashMap};
18use std::sync::{Arc, Mutex};
19
20use crate::ast::PortType;
21use crate::iteration::comprehension::flatten::flatten_static_sources;
22use crate::iteration::comprehension::source::{LiteralValue, Source};
23use crate::iteration::comprehension::{
24 Comprehension, Mode as ValidationMode, StreamerValue, ValidationWarning,
25};
26use crate::kernel::PolydatProgram;
27use crate::kernel::interp::{Lookup, NoScope};
28
29use super::ast::{
30 Arg, Binding, BindingModifier, CallExpr, Expr, ExternPort, ForSource, ForSourceKind, ForStmt,
31 InputDecl, PolydatFile, Statement,
32};
33use super::lexer::Span;
34
35#[derive(Debug, Clone)]
38pub struct Traversal {
39 pub span: Span,
41 pub source_text: String,
43 pub comprehension: Comprehension,
46 pub elements: Vec<(String, PortType)>,
48 pub cascade: Vec<(String, PortType)>,
51 pub program: Arc<PolydatProgram>,
54 pub body: Arc<BodySource>,
58}
59
60pub struct BodySource {
65 pub(crate) file: PolydatFile,
66 pub(crate) source_text: String,
67 pub(crate) source_dir: Option<std::path::PathBuf>,
68 pub(crate) lib_paths: Vec<std::path::PathBuf>,
69 pub(crate) strict: bool,
70 pub(crate) context_label: String,
71 pub(crate) cursor_limit: Option<u64>,
72 pub(crate) pragmas: super::pragmas::PragmaSet,
73 pub(super) modules: HashMap<String, super::modules::ResolvedModule>,
77 pub(crate) programs: Mutex<HashMap<crate::Engine, Arc<dyn crate::kernel::KernelProgram>>>,
79 pub(crate) ledger: Arc<crate::kernel::CompileLedger>,
82 pub(crate) resources: crate::resource::ResourceScope,
85}
86
87impl BodySource {
88 pub fn source_text(&self) -> &str {
92 &self.source_text
93 }
94
95 pub fn program_on(
104 &self,
105 engine: crate::Engine,
106 ) -> Result<Arc<dyn crate::kernel::KernelProgram>, crate::KernelError> {
107 let mut programs = self
108 .programs
109 .lock()
110 .unwrap_or_else(|poisoned| poisoned.into_inner());
111 if let Some(program) = programs.get(&engine) {
112 return Ok(program.clone());
113 }
114 let program = super::compile::Compiler::compile_body_on(self, engine)?.into_program();
115 programs.insert(engine, program.clone());
116 Ok(program)
117 }
118
119 #[allow(clippy::too_many_arguments)]
121 pub(super) fn from_parts(
122 file: PolydatFile,
123 source_text: String,
124 source_dir: Option<std::path::PathBuf>,
125 lib_paths: Vec<std::path::PathBuf>,
126 strict: bool,
127 context_label: String,
128 cursor_limit: Option<u64>,
129 pragmas: super::pragmas::PragmaSet,
130 modules: HashMap<String, super::modules::ResolvedModule>,
131 ledger: Arc<crate::kernel::CompileLedger>,
132 resources: crate::resource::ResourceScope,
133 ) -> Self {
134 BodySource {
135 file,
136 source_text,
137 source_dir,
138 lib_paths,
139 strict,
140 context_label,
141 cursor_limit,
142 pragmas,
143 modules,
144 programs: Mutex::new(HashMap::new()),
145 ledger,
146 resources,
147 }
148 }
149
150 pub fn from_source(source: &str, context_label: &str) -> Result<Self, String> {
159 Self::from_source_in(
160 source,
161 context_label,
162 crate::kernel::CompileLedger::new(),
163 crate::resource::ResourceScope::new(),
164 )
165 }
166
167 pub(crate) fn from_source_in(
171 source: &str,
172 context_label: &str,
173 ledger: Arc<crate::kernel::CompileLedger>,
174 resources: crate::resource::ResourceScope,
175 ) -> Result<Self, String> {
176 let tokens = super::lexer::lex(source)?;
177 let file = super::parser::parse(tokens)?;
178 Ok(BodySource {
179 file,
180 source_text: source.to_string(),
181 source_dir: None,
182 lib_paths: Vec::new(),
183 strict: false,
184 context_label: context_label.to_string(),
185 cursor_limit: None,
186 pragmas: super::pragmas::PragmaSet::default(),
187 modules: HashMap::new(),
188 programs: Mutex::new(HashMap::new()),
189 ledger,
190 resources,
191 })
192 }
193}
194
195impl std::fmt::Debug for BodySource {
196 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
197 f.debug_struct("BodySource")
198 .field("context", &self.context_label)
199 .field("statements", &self.file.statements.len())
200 .finish()
201 }
202}
203
204impl Traversal {
205 pub fn program_on(
210 &self,
211 engine: crate::Engine,
212 ) -> Result<Arc<dyn crate::kernel::KernelProgram>, crate::KernelError> {
213 if matches!(engine, crate::Engine::Interpreter(_)) {
214 return Ok(self.program.clone());
215 }
216 self.body.program_on(engine)
217 }
218}
219
220#[derive(Debug, Clone)]
223pub struct Producer {
224 pub name: String,
226 pub span: Span,
228 pub source_text: String,
230 pub comprehension: Comprehension,
232}
233
234pub fn strip_for_forms(
245 file: &PolydatFile,
246 mode: ValidationMode,
247 scope: &dyn Lookup,
248 events: &mut Vec<super::events::CompileEvent>,
249 has: &dyn Fn(&str) -> bool,
250) -> Result<(PolydatFile, Vec<ForStmt>, Vec<Producer>), String> {
251 let mut parent = Vec::with_capacity(file.statements.len());
252 let mut fors = Vec::new();
253 let mut producers: Vec<Producer> = Vec::new();
254 for stmt in &file.statements {
255 match stmt {
256 Statement::For(f) => fors.push(f.clone()),
257 Statement::Binding(b) if matches!(b.value, Expr::For(_)) => {
258 let Expr::For(source) = &b.value else {
259 unreachable!()
260 };
261 let (comprehension, warnings) =
262 resolve_source_with(source, &producers, mode, scope)?;
263 events.extend(warning_events(source, &warnings));
264 let name = b.targets.join(",");
265 let value = StreamerValue::in_scope(source.to_text(), comprehension.clone(), has);
266 parent.push(Statement::Binding(Binding {
267 targets: b.targets.clone(),
268 value: Expr::Call(CallExpr {
269 func: "streamer".into(),
270 args: vec![Arg::Positional(Expr::StringLit(value.to_json(), b.span))],
271 span: b.span,
272 }),
273 modifier: BindingModifier::CONST,
274 type_annotation: None,
275 span: b.span,
276 }));
277 producers.push(Producer {
278 name,
279 span: b.span,
280 source_text: source.to_text(),
281 comprehension,
282 });
283 }
284 other => parent.push(other.clone()),
285 }
286 }
287 Ok((PolydatFile { statements: parent }, fors, producers))
288}
289
290pub fn resolve_source(source: &ForSource, producers: &[Producer]) -> Result<Comprehension, String> {
294 resolve_source_with(
295 source,
296 producers,
297 ValidationMode::Permissive,
298 &NoScope::new(),
299 )
300 .map(|(c, _)| c)
301}
302
303pub(super) fn warning_events(
306 source: &ForSource,
307 warnings: &[ValidationWarning],
308) -> Vec<super::events::CompileEvent> {
309 warnings
310 .iter()
311 .map(|w| super::events::CompileEvent::ComprehensionWarning {
312 source: source.to_text(),
313 line: source.span.line,
314 col: source.span.col,
315 warning: w.to_string(),
316 })
317 .collect()
318}
319
320pub fn resolve_source_with(
324 source: &ForSource,
325 producers: &[Producer],
326 mode: ValidationMode,
327 scope: &dyn Lookup,
328) -> Result<(Comprehension, Vec<ValidationWarning>), String> {
329 let find = |name: &str| -> Result<Comprehension, String> {
330 producers
331 .iter()
332 .rev()
333 .find(|p| p.name == name)
334 .map(|p| p.comprehension.clone())
335 .ok_or_else(|| {
336 let known: Vec<&str> = producers.iter().map(|p| p.name.as_str()).collect();
337 format!(
338 "`for {}` at line {}, col {}: no producer named '{name}' is bound in this scope{}",
339 source.to_text(),
340 source.span.line,
341 source.span.col,
342 if known.is_empty() { String::new() } else { format!("; producers here: {}", known.join(", ")) }
343 )
344 })
345 };
346 let at = |e: &dyn std::fmt::Display| {
347 format!(
348 "`for {}` at line {}, col {}: {e}",
349 source.to_text(),
350 source.span.line,
351 source.span.col
352 )
353 };
354 let comprehension = match &source.kind {
355 ForSourceKind::Comprehension(c) => {
356 let written = source.to_text();
364 let from_text =
365 crate::iteration::comprehension::spec::parse_comprehension_algebra(&written)
366 .map_err(|e| at(&e))?;
367 if from_text != *c {
368 return Err(at(&format!(
369 "this comprehension does not survive being written and read back: it writes \
370 as `{written}`, which reads as a different comprehension. That is an \
371 inconsistency between the comprehension renderer and parser, not an error \
372 in this program"
373 )));
374 }
375 c.clone()
376 }
377 ForSourceKind::Producer(name) => find(name)?,
378 ForSourceKind::Derived {
379 base,
380 filter,
381 order,
382 } => {
383 let mut c = find(base)?;
384 if let Some(pred) = filter {
385 c = Comprehension::filter(c, pred.clone());
386 }
387 if let Some(spec) = order {
388 let (strategy, truncation, seed) = parse_order(spec).map_err(|e| at(&e))?;
389 c = Comprehension::order_seeded(c, strategy, truncation, seed);
390 }
391 c
392 }
393 };
394 let comprehension = flatten_static_sources(&comprehension, scope);
398 if let Some((_, message)) =
401 crate::iteration::comprehension::flatten::first_refused_generator(&comprehension, scope)
402 {
403 return Err(at(&message));
404 }
405 let report =
409 crate::iteration::comprehension::validate(&comprehension, mode).map_err(|e| at(&e))?;
410 Ok((comprehension, report.warnings))
411}
412
413fn parse_order(
416 spec: &str,
417) -> Result<
418 (
419 crate::iteration::comprehension::StrategyName,
420 Option<u64>,
421 Option<u64>,
422 ),
423 String,
424> {
425 crate::iteration::comprehension::spec::parse_order(spec)
426}
427
428pub fn element_types(
432 comprehension: &Comprehension,
433 probe: &mut dyn FnMut(&str) -> Result<PortType, String>,
434) -> Result<Vec<(String, PortType)>, String> {
435 let mut out = Vec::new();
436 collect_element_types(comprehension, probe, &mut out)?;
437 Ok(out)
438}
439
440fn collect_element_types(
441 c: &Comprehension,
442 probe: &mut dyn FnMut(&str) -> Result<PortType, String>,
443 out: &mut Vec<(String, PortType)>,
444) -> Result<(), String> {
445 match c {
446 Comprehension::Clause { name, source } => {
447 if out.iter().any(|(n, _)| n == name) {
448 return Ok(());
449 }
450 let ty = source_type(name, source, probe)?;
451 out.push((name.clone(), ty));
452 Ok(())
453 }
454 Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
455 for child in children {
456 collect_element_types(child, probe, out)?;
457 }
458 Ok(())
459 }
460 Comprehension::Union { children } => {
461 if let Some(first) = children.first() {
464 collect_element_types(first, probe, out)?;
465 }
466 Ok(())
467 }
468 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
469 collect_element_types(child, probe, out)
470 }
471 }
472}
473
474fn source_type(
475 name: &str,
476 source: &Source,
477 probe: &mut dyn FnMut(&str) -> Result<PortType, String>,
478) -> Result<PortType, String> {
479 match source {
480 Source::Literal { values } => {
481 let mut ty: Option<PortType> = None;
482 for v in values {
483 let t = match v {
484 LiteralValue::Int(_) | LiteralValue::UInt(_) => PortType::U64,
485 LiteralValue::Float(_) => PortType::F64,
486 LiteralValue::String(_) => PortType::Str,
487 LiteralValue::Bool(_) => PortType::Bool,
488 LiteralValue::Json(_) => PortType::Json,
489 };
490 match ty {
491 None => ty = Some(t),
492 Some(PortType::F64) if t == PortType::U64 => {}
494 Some(PortType::U64) if t == PortType::F64 => ty = Some(PortType::F64),
495 Some(prev) if prev != t => {
496 return Err(format!(
497 "element '{name}': literal list mixes {prev:?} and {t:?} values; a comprehension element has one type"
498 ));
499 }
500 Some(_) => {}
501 }
502 }
503 ty.ok_or_else(|| format!("element '{name}': literal list is empty"))
504 }
505 Source::IntRange { .. } => Ok(PortType::U64),
506 Source::ContinuousInterval { .. } | Source::Distribution { .. } => Ok(PortType::F64),
507 Source::WorkloadParamList { .. } => Ok(PortType::Str),
509 Source::Generator { expr, .. } => {
510 let head = expr.trim();
511 if head.starts_with("partitions(")
512 || head.starts_with("subdivide(")
513 || head.ends_with(".partitions")
514 {
515 return Ok(PortType::Ext);
516 }
517 if let Some(g) = crate::iteration::comprehension::eval::NamedGenerator::of_call(head) {
520 return Ok(if g.yields_integers() {
521 PortType::U64
522 } else {
523 PortType::F64
524 });
525 }
526 probe(head)
527 .map_err(|e| format!("element '{name}': cannot type generator `{head}`: {e}"))
528 }
529 }
530}
531
532fn body_declared(body: &[Statement], elements: &[(String, PortType)]) -> BTreeSet<String> {
536 let mut names: BTreeSet<String> = elements.iter().map(|(n, _)| n.clone()).collect();
537 names.insert("cycle".to_string());
538 for stmt in body {
539 match stmt {
540 Statement::Binding(b) => names.extend(b.targets.iter().cloned()),
541 Statement::InputDecl(d) => {
542 names.insert(d.name.clone());
543 }
544 Statement::ExternPort(p) => {
545 names.insert(p.name.clone());
546 }
547 Statement::ModuleDef(m) => {
548 names.insert(m.name.clone());
549 }
550 Statement::Cursor(c) => {
551 names.insert(c.name.clone());
552 }
553 Statement::Pragma { .. } => {}
554 Statement::For(f) => {
555 let _ = f;
558 }
559 Statement::Tile(t) => {
560 names.insert(t.name.clone());
561 }
562 }
563 }
564 names
565}
566
567fn tile_references(pieces: &[super::ast::TilePiece], out: &mut BTreeSet<String>) {
569 use super::ast::TilePiece;
570 use super::refs::collect_expr_refs;
571 for piece in pieces {
572 match piece {
573 TilePiece::Static(_) => {}
574 TilePiece::Hole(h) => collect_expr_refs(&h.expr, out),
575 TilePiece::Projection { body, .. } => tile_references(body, out),
576 TilePiece::Branch {
577 cond,
578 then,
579 otherwise,
580 ..
581 } => {
582 collect_expr_refs(cond, out);
583 tile_references(then, out);
584 if let Some(o) = otherwise {
585 tile_references(o, out);
586 }
587 }
588 }
589 }
590}
591
592fn body_references(body: &[Statement], out: &mut BTreeSet<String>) {
597 use super::refs::collect_expr_refs;
598 for stmt in body {
599 match stmt {
600 Statement::Binding(b) => collect_expr_refs(&b.value, out),
601 Statement::ExternPort(p) => {
602 if let Some(d) = &p.default {
603 collect_expr_refs(d, out);
604 }
605 }
606 Statement::Cursor(c) => {
607 collect_expr_refs(&c.constructor, out);
608 if let Some(over) = &c.over {
609 collect_expr_refs(over, out);
610 }
611 }
612 Statement::For(f) => {
613 let mut inner = BTreeSet::new();
614 body_references(&f.body, &mut inner);
615 let own = body_declared(&f.body, &[]);
616 let elems: BTreeSet<String> = f.source.element_names().into_iter().collect();
617 for n in inner {
618 if !own.contains(&n) && !elems.contains(&n) {
619 out.insert(n);
620 }
621 }
622 }
623 Statement::Tile(t) => tile_references(&t.pieces, out),
624 Statement::InputDecl(_) | Statement::ModuleDef(_) | Statement::Pragma { .. } => {}
625 }
626 }
627}
628
629pub fn child_file(
634 f: &ForStmt,
635 comprehension: &Comprehension,
636 elements: &[(String, PortType)],
637 type_of: &dyn Fn(&str) -> Option<PortType>,
638) -> Result<(PolydatFile, Vec<(String, PortType)>), String> {
639 for stmt in &f.body {
641 if let Statement::InputDecl(d) = stmt
642 && d.name != "cycle"
643 {
644 return Err(format!(
645 "`for {}` at line {}, col {}: a traversal body cannot declare input '{}'; only `cycle` is a coordinate inside a body, and the comprehension supplies the rest",
646 f.source.to_text(),
647 f.span.line,
648 f.span.col,
649 d.name
650 ));
651 }
652 }
653 let declared = body_declared(&f.body, elements);
654 let mut referenced = BTreeSet::new();
655 body_references(&f.body, &mut referenced);
656 referenced.extend(comprehension.source_names_read());
659
660 let mut cascade = Vec::new();
661 for name in referenced {
662 if declared.contains(&name) {
663 continue;
664 }
665 let ty = type_of(&name);
666 if let Some(ty) = ty {
667 cascade.push((name, ty));
668 }
669 }
672
673 let span = f.span;
674 let mut statements = Vec::with_capacity(f.body.len() + elements.len() + cascade.len() + 1);
675 if !f.body.iter().any(|s| matches!(s, Statement::InputDecl(_))) {
676 statements.push(Statement::InputDecl(InputDecl {
677 name: "cycle".into(),
678 ty: Some("u64".into()),
679 span,
680 }));
681 }
682 for (name, ty) in elements {
683 statements.push(Statement::ExternPort(ExternPort {
684 name: name.clone(),
685 typ: ty.to_keyword().to_string(),
686 default: None,
687 span,
688 }));
689 }
690 for (name, ty) in &cascade {
691 statements.push(Statement::ExternPort(ExternPort {
692 name: name.clone(),
693 typ: ty.to_keyword().to_string(),
694 default: None,
695 span,
696 }));
697 }
698 statements.extend(f.body.iter().cloned());
699 Ok((PolydatFile { statements }, cascade))
700}