1use bumpalo::{Bump, collections::Vec as BumpVec};
2use itertools::zip_eq;
3use rustc_hash::FxHashMap;
4use thiserror::Error;
5
6use super::{
7 LinkName, Module, Node, Operation, Param, Region, SeqPart, Symbol, SymbolIdent, SymbolName,
8 Term, VarName,
9};
10use crate::v0::{RegionKind, ScopeClosure, table};
11use crate::v0::{
12 scope::{LinkTable, SymbolTable, VarTable},
13 table::{LinkIndex, NodeId, RegionId, TermId, VarId},
14};
15
16pub struct Context<'a> {
17 module: table::Module<'a>,
18 bump: &'a Bump,
19 vars: VarTable<'a>,
20 links: LinkTable<&'a str>,
21 symbols: SymbolTable<'a>,
22 imports: FxHashMap<SymbolIdent, NodeId>,
23 terms: FxHashMap<table::Term<'a>, TermId>,
24}
25
26impl<'a> Context<'a> {
27 pub fn new(bump: &'a Bump) -> Self {
37 Self {
38 module: table::Module::default(),
39 bump,
40 vars: VarTable::new(),
41 links: LinkTable::new(),
42 symbols: SymbolTable::new(),
43 imports: FxHashMap::default(),
44 terms: FxHashMap::default(),
45 }
46 }
47
48 pub fn resolve_module(&mut self, module: &'a Module) -> BuildResult<()> {
49 self.module.root = self.module.insert_region(table::Region::default());
50 self.symbols.enter(self.module.root);
51 self.links.enter(self.module.root);
52
53 let children = self.resolve_nodes(&module.root.children)?;
54 let meta = self.resolve_terms(&module.root.meta)?;
55
56 let (links, ports) = self.links.exit();
57 self.symbols.exit();
58 let scope = Some(table::RegionScope { links, ports });
59
60 let all_children = {
63 let mut all_children =
64 BumpVec::with_capacity_in(children.len() + self.imports.len(), self.bump);
65 all_children.extend(children);
66 all_children.extend(self.imports.drain().map(|(_, node)| node));
67 all_children.into_bump_slice()
68 };
69
70 self.module.regions[self.module.root.index()] = table::Region {
71 kind: RegionKind::Module,
72 sources: &[],
73 targets: &[],
74 children: all_children,
75 meta,
76 signature: None,
77 scope,
78 };
79
80 Ok(())
81 }
82
83 fn resolve_terms(&mut self, terms: &'a [Term]) -> BuildResult<&'a [TermId]> {
84 try_alloc_slice(self.bump, terms.iter().map(|term| self.resolve_term(term)))
85 }
86
87 fn resolve_term(&mut self, term: &'a Term) -> BuildResult<TermId> {
88 let term = match term {
89 Term::Wildcard => table::Term::Wildcard,
90 Term::Var(var_name) => table::Term::Var(self.resolve_var(var_name)?),
91 Term::Apply(symbol_ident, terms) => {
92 let symbol_id = self.resolve_symbol_ident(symbol_ident);
93 let terms = self.resolve_terms(terms)?;
94 table::Term::Apply(symbol_id, terms)
95 }
96 Term::List(parts) => table::Term::List(self.resolve_seq_parts(parts)?),
97 Term::Literal(literal) => table::Term::Literal(literal.clone()),
98 Term::Tuple(parts) => table::Term::Tuple(self.resolve_seq_parts(parts)?),
99 Term::Func(region) => {
100 let region = self.resolve_region(region, ScopeClosure::Closed)?;
101 table::Term::Func(region)
102 }
103 };
104
105 Ok(*self
106 .terms
107 .entry(term.clone())
108 .or_insert_with(|| self.module.insert_term(term)))
109 }
110
111 fn resolve_seq_parts(&mut self, parts: &'a [SeqPart]) -> BuildResult<&'a [table::SeqPart]> {
112 try_alloc_slice(
113 self.bump,
114 parts.iter().map(|part| self.resolve_seq_part(part)),
115 )
116 }
117
118 fn resolve_seq_part(&mut self, part: &'a SeqPart) -> BuildResult<table::SeqPart> {
119 Ok(match part {
120 SeqPart::Item(term) => table::SeqPart::Item(self.resolve_term(term)?),
121 SeqPart::Splice(term) => table::SeqPart::Splice(self.resolve_term(term)?),
122 })
123 }
124
125 fn resolve_nodes(&mut self, nodes: &'a [Node]) -> BuildResult<&'a [NodeId]> {
126 let ids: &[_] = self.bump.alloc_slice_fill_with(nodes.len(), |_| {
128 self.module.insert_node(table::Node::default())
129 });
130
131 for (id, node) in zip_eq(ids, nodes) {
136 if let Some((symbol_name, version)) = node.operation.symbol_binding() {
137 match &node.operation {
138 Operation::Import(_) => {
139 self.symbols
140 .insert_import(symbol_name.as_ref(), version, *id)
141 }
142 _ => self.symbols.insert(symbol_name.as_ref(), version, *id),
143 }
144 .map_err(|_| ResolveError::DuplicateSymbol(symbol_name.clone()))?;
145 }
146 }
147
148 for (id, node) in zip_eq(ids, nodes) {
150 self.resolve_node(*id, node)?;
151 }
152
153 Ok(ids)
154 }
155
156 fn resolve_node(&mut self, node_id: NodeId, node: &'a Node) -> BuildResult<()> {
157 let inputs = self.resolve_links(&node.inputs)?;
158 let outputs = self.resolve_links(&node.outputs)?;
159
160 if node.operation.symbol().is_some() {
162 self.vars.enter(node_id);
163 }
164
165 let mut scope_closure = ScopeClosure::Open;
166
167 let operation = match &node.operation {
168 Operation::Invalid => table::Operation::Invalid,
169 Operation::Dfg => table::Operation::Dfg,
170 Operation::Cfg => table::Operation::Cfg,
171 Operation::Block => table::Operation::Block,
172 Operation::TailLoop => table::Operation::TailLoop,
173 Operation::Conditional => table::Operation::Conditional,
174 Operation::DefineFunc(symbol) => {
175 let symbol = self.resolve_symbol(symbol)?;
176 scope_closure = ScopeClosure::Closed;
177 table::Operation::DefineFunc(symbol)
178 }
179 Operation::DeclareFunc(symbol) => {
180 let symbol = self.resolve_symbol(symbol)?;
181 table::Operation::DeclareFunc(symbol)
182 }
183 Operation::DefineAlias(symbol, term) => {
184 let symbol = self.resolve_symbol(symbol)?;
185 let term = self.resolve_term(term)?;
186 table::Operation::DefineAlias(symbol, term)
187 }
188 Operation::DeclareAlias(symbol) => {
189 let symbol = self.resolve_symbol(symbol)?;
190 table::Operation::DeclareAlias(symbol)
191 }
192 Operation::DeclareConstructor(symbol) => {
193 let symbol = self.resolve_symbol(symbol)?;
194 table::Operation::DeclareConstructor(symbol)
195 }
196 Operation::DeclareOperation(symbol) => {
197 let symbol = self.resolve_symbol(symbol)?;
198 table::Operation::DeclareOperation(symbol)
199 }
200 Operation::Import(symbol_ident) => table::Operation::Import {
201 name: symbol_ident.name.as_ref(),
202 version: &symbol_ident.version,
203 },
204 Operation::Custom(term) => {
205 let term = self.resolve_term(term)?;
206 table::Operation::Custom(term)
207 }
208 };
209
210 let meta = self.resolve_terms(&node.meta)?;
211 let regions = self.resolve_regions(&node.regions, scope_closure)?;
212
213 let signature = match &node.signature {
214 Some(signature) => Some(self.resolve_term(signature)?),
215 None => None,
216 };
217
218 if node.operation.symbol().is_some() {
220 self.vars.exit();
221 }
222
223 self.module.nodes[node_id.index()] = table::Node {
224 operation,
225 inputs,
226 outputs,
227 regions,
228 meta,
229 signature,
230 };
231
232 Ok(())
233 }
234
235 fn resolve_links(&mut self, links: &'a [LinkName]) -> BuildResult<&'a [LinkIndex]> {
236 try_alloc_slice(self.bump, links.iter().map(|link| self.resolve_link(link)))
237 }
238
239 fn resolve_link(&mut self, link: &'a LinkName) -> BuildResult<LinkIndex> {
240 Ok(self.links.use_link(link.as_ref()))
241 }
242
243 fn resolve_regions(
244 &mut self,
245 regions: &'a [Region],
246 scope_closure: ScopeClosure,
247 ) -> BuildResult<&'a [RegionId]> {
248 try_alloc_slice(
249 self.bump,
250 regions
251 .iter()
252 .map(|region| self.resolve_region(region, scope_closure)),
253 )
254 }
255
256 fn resolve_region(
257 &mut self,
258 region: &'a Region,
259 scope_closure: ScopeClosure,
260 ) -> BuildResult<RegionId> {
261 let meta = self.resolve_terms(®ion.meta)?;
262 let signature = match ®ion.signature {
263 Some(signature) => Some(self.resolve_term(signature)?),
264 None => None,
265 };
266
267 let region_id = self.module.insert_region(table::Region::default());
270
271 self.symbols.enter(region_id);
273
274 if ScopeClosure::Closed == scope_closure {
276 self.links.enter(region_id);
277 }
278
279 let sources = self.resolve_links(®ion.sources)?;
280 let targets = self.resolve_links(®ion.targets)?;
281 let children = self.resolve_nodes(®ion.children)?;
282
283 let scope = match scope_closure {
285 ScopeClosure::Open => None,
286 ScopeClosure::Closed => {
287 let (links, ports) = self.links.exit();
288 Some(table::RegionScope { links, ports })
289 }
290 };
291 self.symbols.exit();
292
293 self.module.regions[region_id.index()] = table::Region {
294 kind: region.kind,
295 sources,
296 targets,
297 children,
298 meta,
299 signature,
300 scope,
301 };
302
303 Ok(region_id)
304 }
305
306 fn resolve_symbol(&mut self, symbol: &'a Symbol) -> BuildResult<&'a table::Symbol<'a>> {
307 let name = symbol.name.as_ref();
308 let visibility = &symbol.visibility;
309 let params = self.resolve_params(&symbol.params)?;
310 let constraints = self.resolve_terms(&symbol.constraints)?;
311 let signature = self.resolve_term(&symbol.signature)?;
312
313 Ok(self.bump.alloc(table::Symbol {
314 visibility,
315 name,
316 version: &symbol.version,
317 params,
318 constraints,
319 signature,
320 }))
321 }
322
323 fn resolve_params(&mut self, params: &'a [Param]) -> BuildResult<&'a [table::Param<'a>]> {
329 try_alloc_slice(
330 self.bump,
331 params.iter().map(|param| self.resolve_param(param)),
332 )
333 }
334
335 fn resolve_param(&mut self, param: &'a Param) -> BuildResult<table::Param<'a>> {
340 let name = param.name.as_ref();
341 let r#type = self.resolve_term(¶m.r#type)?;
342
343 self.vars
344 .insert(param.name.as_ref())
345 .map_err(|_| ResolveError::DuplicateVar(param.name.clone()))?;
346
347 Ok(table::Param { name, r#type })
348 }
349
350 fn resolve_var(&self, var_name: &'a VarName) -> BuildResult<VarId> {
351 self.vars
352 .resolve(var_name.as_ref())
353 .map_err(|_| ResolveError::UnknownVar(var_name.clone()))
354 }
355
356 fn resolve_symbol_ident(&mut self, symbol_ident: &'a SymbolIdent) -> NodeId {
373 if let Ok(node) = self
374 .symbols
375 .resolve(symbol_ident.name.as_ref(), symbol_ident.version.as_ref())
376 {
377 return node;
378 }
379
380 *self.imports.entry(symbol_ident.clone()).or_insert_with(|| {
381 self.module.insert_node(table::Node {
382 operation: table::Operation::Import {
383 name: symbol_ident.name.as_ref(),
384 version: &symbol_ident.version,
385 },
386 ..Default::default()
387 })
388 })
389 }
390
391 pub fn finish(self) -> table::Module<'a> {
392 self.module
393 }
394}
395
396#[derive(Debug, Clone, Error)]
398#[non_exhaustive]
399#[error("Error resolving model module")]
400pub enum ResolveError {
401 #[error("unknown var: {0}")]
403 UnknownVar(VarName),
404 #[error("duplicate var: {0}")]
406 DuplicateVar(VarName),
407 #[error("duplicate symbol: {0}")]
409 DuplicateSymbol(SymbolName),
410}
411
412type BuildResult<T> = Result<T, ResolveError>;
413
414fn try_alloc_slice<T, E>(
415 bump: &Bump,
416 iter: impl IntoIterator<Item = Result<T, E>>,
417) -> Result<&[T], E> {
418 let iter = iter.into_iter();
419 let mut vec = BumpVec::with_capacity_in(iter.size_hint().0, bump);
420 for item in iter {
421 vec.push(item?);
422 }
423 Ok(vec.into_bump_slice())
424}
425
426#[cfg(test)]
427mod test {
428 use crate::v0::{ast, table};
429 use bumpalo::Bump;
430 use rstest::rstest;
431 use std::str::FromStr as _;
432
433 #[derive(Debug, Clone, Copy)]
434 enum SymbolUse {
435 Custom,
436 Meta(usize),
437 }
438
439 #[rstest]
440 fn root_var_errors() {
441 let text = "(hugr 0) (mod) (meta ?x)";
442 let ast = ast::Package::from_str(text.trim()).unwrap();
443 assert!(ast.resolve(&Bump::new()).is_err());
444 }
445
446 #[rstest]
448 #[case::operation(
449 "
450 (hugr 0)
451 (mod)
452 (import someOp@0.2.3)
453 (someOp)
454 (declare-operation someOp@0.3.0 (core.fn [] []))
455 ",
456 SymbolUse::Custom,
457 "someOp",
458 Some("0.3.0")
459 )]
460 #[case::type_term(
461 "
462 (hugr 0)
463 (mod)
464 (meta someType)
465 (import someType@0.2.3)
466 ",
467 SymbolUse::Meta(0),
468 "someType",
469 Some("0.2.3")
470 )]
471 #[case::type_latest(
472 "
473 (hugr 0)
474 (mod)
475 (meta someType)
476 (import someType@0.2.3)
477 (import someType@0.3.0)
478 ",
479 SymbolUse::Meta(0),
480 "someType",
481 Some("0.3.0")
482 )]
483 #[case::unversioned_import(
484 "
485 (hugr 0)
486 (mod)
487 (meta someType)
488 (import someType)
489 ",
490 SymbolUse::Meta(0),
491 "someType",
492 None
493 )]
494 #[case::unversioned_operation(
495 "
496 (hugr 0)
497 (mod)
498 (someOp)
499 (declare-operation someOp (core.fn [] []))
500 ",
501 SymbolUse::Custom,
502 "someOp",
503 None
504 )]
505 fn unversioned_resolution(
506 #[case] text: &str,
507 #[case] symbol_use: SymbolUse,
508 #[case] expected_name: &str,
509 #[case] expected_version: Option<&str>,
510 ) {
511 let bump = Bump::new();
512 let ast = ast::Package::from_str(text.trim()).unwrap();
513 let package = ast.resolve(&bump).unwrap();
514 let module = &package.modules[0];
515
516 assert_symbol_use(module, symbol_use, expected_name, expected_version);
517 }
518
519 fn assert_symbol_use(
520 module: &table::Module<'_>,
521 symbol_use: SymbolUse,
522 expected_name: &str,
523 expected_version: Option<&str>,
524 ) {
525 let term = match symbol_use {
526 SymbolUse::Custom => custom_term(module),
527 SymbolUse::Meta(index) => module.get_region(module.root).unwrap().meta[index],
528 };
529
530 assert_symbol_application(module, term, expected_name, expected_version);
531 }
532
533 fn custom_term(module: &table::Module<'_>) -> table::TermId {
534 let root = module.get_region(module.root).unwrap();
535 root.children
536 .iter()
537 .find_map(|node_id| {
538 let node = module.get_node(*node_id).unwrap();
539 match node.operation {
540 table::Operation::Custom(term) => Some(term),
541 _ => None,
542 }
543 })
544 .expect("expected a custom operation node")
545 }
546
547 fn assert_symbol_application(
548 module: &table::Module<'_>,
549 term: table::TermId,
550 expected_name: &str,
551 expected_version: Option<&str>,
552 ) {
553 let table::Term::Apply(symbol, args) = module.get_term(term).unwrap() else {
554 panic!("expected term to be a symbol application");
555 };
556 assert!(args.is_empty());
557
558 let node = module.get_node(*symbol).unwrap();
559 let name = node.operation.symbol().expect("expected symbol node");
560 let version = node
561 .operation
562 .symbol_version()
563 .expect("expected symbol version");
564 assert_eq!(name, expected_name);
565 assert_eq!(
566 version.as_ref().map(ToString::to_string).as_deref(),
567 expected_version
568 );
569 }
570}