Skip to main content

hugr_model/v0/ast/
resolve.rs

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    /// Create an empty resolver context backed by `bump`.
28    ///
29    /// The context records explicit symbol versions from imports and
30    /// declarations. It does not fill missing versions from any external
31    /// registry; unresolved legacy versions remain `None` in the table model.
32    ///
33    /// TODO: Remove the `None` compatibility path once encoded HUGRs without
34    /// extension version information are no longer supported.
35    /// <http://github.com/Quantinuum/hugr/issues/3086>
36    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        // Symbols that could not be resolved within the module still need to
61        // be represented by a node. This is why we add import nodes.
62        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        // Allocate ids for all nodes by introducing placeholders into the module.
127        let ids: &[_] = self.bump.alloc_slice_fill_with(nodes.len(), |_| {
128            self.module.insert_node(table::Node::default())
129        });
130
131        // For those nodes that introduce symbols, we then associate the symbol
132        // with the id of the node. This serves as a form of forward declaration
133        // so that the symbol is visible in the current region regardless of the
134        // order of the nodes.
135        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        // Finally we can build the actual nodes.
149        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        // When the node introduces a symbol it also introduces a new variable scope.
161        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        // We need to close the variable scope if we have opened one before.
219        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(&region.meta)?;
262        let signature = match &region.signature {
263            Some(signature) => Some(self.resolve_term(signature)?),
264            None => None,
265        };
266
267        // We insert a placeholder for the region in order to allocate a region
268        // id, which we need to track the region's scopes.
269        let region_id = self.module.insert_region(table::Region::default());
270
271        // Each region defines a new scope for symbols.
272        self.symbols.enter(region_id);
273
274        // If the region is closed, it also defines a new scope for links.
275        if ScopeClosure::Closed == scope_closure {
276            self.links.enter(region_id);
277        }
278
279        let sources = self.resolve_links(&region.sources)?;
280        let targets = self.resolve_links(&region.targets)?;
281        let children = self.resolve_nodes(&region.children)?;
282
283        // Close the region's scopes.
284        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    /// Builds symbol parameters.
324    ///
325    /// This incrementally inserts the names of the parameters into the current
326    /// variable scope, so that any parameter is in scope for each of its
327    /// succeeding parameters.
328    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    /// Builds a symbol parameter.
336    ///
337    /// This inserts the name of the parameter into the current variable scope,
338    /// making the parameter accessible as a variable.
339    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(&param.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    /// Resolves a symbol identifier and returns the node that introduces it.
357    ///
358    /// Versioned references resolve exactly. Unversioned references first use
359    /// any unversioned binding, then the latest visible versioned binding.
360    /// If no binding exists, the implicitly created import preserves the
361    /// reference version as written, including `None` for legacy unversioned
362    /// references.
363    ///
364    /// TODO: Remove the legacy unversioned fallback once encoded HUGRs without
365    /// extension version information are no longer supported.
366    /// <http://github.com/Quantinuum/hugr/issues/3086>
367    ///
368    /// When there is no symbol with this name in scope, we create a new import
369    /// node in the module and record that the symbol has been implicitly
370    /// imported. At the end of the building process, these import nodes are
371    /// inserted into the module's scope.
372    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/// Error that may occur in [`Module::resolve`].
397#[derive(Debug, Clone, Error)]
398#[non_exhaustive]
399#[error("Error resolving model module")]
400pub enum ResolveError {
401    /// Unknown variable.
402    #[error("unknown var: {0}")]
403    UnknownVar(VarName),
404    /// Duplicate variable definition in the same symbol.
405    #[error("duplicate var: {0}")]
406    DuplicateVar(VarName),
407    /// Duplicate symbol definition in the same region.
408    #[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    /// Unversioned extension uses resolve to declarations in scope.
447    #[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}