diff --git a/CHANGELOG.md b/CHANGELOG.md index 110765bb..5cd893cf 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -19,6 +19,7 @@ Types of changes: - Added an `include_dir` kwarg to `loads()` and `load()`, naming the directory custom `include` statements resolve against. A program given as a string could not resolve includes at all, and failed later naming the gate rather than the include. Resolution is opt-in: without the kwarg, no files are read. ([#368](https://github.com/qBraid/pyqasm/issues/368)) ### Improved / Modified +- Reduced the length of `visitor.py` by removing the `_handle_function_init_expression` function and adding the `check_only_return_empty` decorator for functions which can be given simple boiler-plate logic for the `self._check_only` parameter. ([#348](https://github.com/qBraid/pyqasm/pull/348)) ### Deprecated diff --git a/src/pyqasm/pulse/validator.py b/src/pyqasm/pulse/validator.py index 888f8180..be8a90bd 100644 --- a/src/pyqasm/pulse/validator.py +++ b/src/pyqasm/pulse/validator.py @@ -114,11 +114,11 @@ def validate_duration_or_stretch_statements( Generic validation function for DurationType and StretchType declarations or assignments. Args: - statement: The AST statement node - statement_type: The expected AST node type - base_type: The declared type (DurationType or StretchType) - rvalue: The initializer or assigned value - global_scope: Global symbol table. + statement (Statement): The AST statement node. + base_type (Any): The declared type, function does nothing + if not DurationType or StretchType. + rvalue (Any): The initializer or assigned value. + global_scope (dict): Global symbol table. Raises: ValidationError: If the assigned value is not a DurationLiteral, diff --git a/src/pyqasm/visitor.py b/src/pyqasm/visitor.py index bd45db75..91380938 100644 --- a/src/pyqasm/visitor.py +++ b/src/pyqasm/visitor.py @@ -20,13 +20,14 @@ """ import copy +import functools import logging import re import sys from collections import OrderedDict, deque from functools import partial from io import StringIO -from typing import Any, Callable, Optional, Sequence, cast +from typing import Any, Callable, Optional, Sequence, TypeVar, cast import numpy as np import openqasm3.ast as qasm3_ast @@ -86,6 +87,22 @@ logger = logging.getLogger(__name__) logger.propagate = False +F = TypeVar("F", bound=Callable[..., Any]) + + +def check_only_return_empty(func: F) -> F: + """Decorator for functions which use check_only to return an empty list.""" + + @functools.wraps(func) + def wrapper(self, *args, **kwargs): + """Wrapper that intercepts the return value and replaces it with an empty list.""" + result = func(self, *args, **kwargs) + if self._check_only: + return [] + return result + + return wrapper # type: ignore + # pylint: disable-next=too-many-instance-attributes class QasmVisitor: @@ -208,6 +225,7 @@ def _construct_visit_map(self): qasm3_ast.CalibrationGrammarDeclaration: self._visit_calibration_grammar_declaration, } + @check_only_return_empty def _visit_quantum_register( self, register: qasm3_ast.QubitDeclaration ) -> list[qasm3_ast.QubitDeclaration]: @@ -289,8 +307,6 @@ def _visit_quantum_register( logger.debug("Added labels for register '%s'", str(register)) - if self._check_only: - return [] return [register] # pylint: disable-next=too-many-locals,too-many-branches,too-many-statements @@ -305,8 +321,7 @@ def _get_op_bits( operation (Any): The operation to get qubits for. qubits (bool): Whether the bits are quantum bits or classical bits. Defaults to True. Returns: - list[IndexedIdentifier | Identifier]: The quantum or classical bits for the operation, - or an empty list if check_only is true. + list[IndexedIdentifier | Identifier]: The quantum or classical bits for the operation. """ openqasm_bits: list[qasm3_ast.IndexedIdentifier | qasm3_ast.Identifier] = [] bit_list = [] @@ -568,26 +583,6 @@ def _qubit_register_consolidation( return _valid_statements - def _handle_function_init_expression( - self, expression: qasm3_ast.FunctionCall, init_value: Any - ) -> None | qasm3_ast.Expression: - """Handle function initialization expression. - - Args: - expression (FunctionCall): The statement to handle function initialization expression. - init_value (Any): The value to handle function initialization expression. - - Returns: - None | Expression: The resultant expression if - the expression is applied, otherwise None. - """ - if isinstance(expression, qasm3_ast.FunctionCall): - func_name = expression.name.name - if func_name in FUNCTION_MAP: - if isinstance(init_value, (float, int)): - return qasm3_ast.FloatLiteral(init_value) - return None - def _handle_extern_function_cleanup( self, statements: list, statement: qasm3_ast.Statement ) -> None: @@ -756,7 +751,6 @@ def _visit_measurement( # pylint: disable=too-many-locals,too-many-branches,too if self._check_only: return [] - return unrolled_measurements def _resolve_unindexed_reset_qubit(self, statement: qasm3_ast.QuantumReset) -> bool: @@ -800,12 +794,13 @@ def _visit_reset(self, statement: qasm3_ast.QuantumReset) -> list[qasm3_ast.Quan or an empty list if self._check_only is True. """ logger.debug("Visiting reset statement '%s'", str(statement)) + if self._resolve_unindexed_reset_qubit(statement): return [statement] - if len(self._function_qreg_size_map) > 0: # atleast in SOME function scope + if len(self._function_qreg_size_map) > 0: # at least in SOME function scope # since we may have multiple function scopes, we need to transform the qubits - # to use the global qreg identifiers + # to use the global qreg identifiers. for transform_map, size_map in zip( reversed(self._function_qreg_transform_map), reversed(self._function_qreg_size_map) ): @@ -845,7 +840,6 @@ def _visit_reset(self, statement: qasm3_ast.QuantumReset) -> list[qasm3_ast.Quan if self._check_only: return [] - return unrolled_resets def _expand_barrier_ranges( @@ -1175,6 +1169,7 @@ def _update_qubit_depth_for_gate( qubit_node.depth = max_involved_depth # pylint: disable=too-many-branches, too-many-locals + @check_only_return_empty def _visit_basic_gate_operation( self, operation: qasm3_ast.QuantumGate, @@ -1273,9 +1268,6 @@ def _visit_basic_gate_operation( for final_gate in result: Qasm3Analyzer.verify_gate_qubits(final_gate, operation.span) - if self._check_only: - return [] - return result def _visit_break(self, statement: qasm3_ast.BreakStatement) -> None: @@ -1308,6 +1300,7 @@ def _is_black_box_gate(self, gate_name: str) -> bool: or gate_name in self._opaque_gates ) + @check_only_return_empty def _visit_custom_gate_operation( self, operation: qasm3_ast.QuantumGate, @@ -1430,11 +1423,9 @@ def _visit_custom_gate_operation( self._scope_manager.pop_scope() self._scope_manager.restore_context() - if self._check_only: - return [] - return result + @check_only_return_empty def _visit_external_gate_operation( self, operation: qasm3_ast.QuantumGate, @@ -1523,11 +1514,10 @@ def gate_function(*qubits): Qasm3Analyzer.verify_gate_qubits(final_gate, operation.span) self._scope_manager.restore_context() - if self._check_only: - return [] return result + @check_only_return_empty def _visit_phase_operation( self, operation: qasm3_ast.QuantumPhase, @@ -1587,9 +1577,6 @@ def _visit_phase_operation( # if it were in function scope, then the args would have been evaluated and added to the # qubit list - if self._check_only: - return [] - return [operation] def _visit_generic_gate_operation( # pylint: disable=too-many-branches, too-many-statements @@ -1759,9 +1746,9 @@ def _neg_x_gates() -> list[qasm3_ast.QuantumGate]: if self._check_only: return [] - return result + @check_only_return_empty def _visit_constant_declaration( self, statement: qasm3_ast.ConstantDeclaration ) -> list[qasm3_ast.Statement]: @@ -1872,18 +1859,16 @@ def _visit_constant_declaration( statement.init_expression = PulseValidator.make_complex_binary_expression(init_value) if isinstance(statement.init_expression, qasm3_ast.FunctionCall): - statement.init_expression = ( - self._handle_function_init_expression(statement.init_expression, init_value) - or statement.init_expression - ) - self._handle_extern_function_cleanup(statements, statement) + function_name = statement.init_expression.name.name + if function_name in FUNCTION_MAP and isinstance(init_value, (float, int)): + statement.init_expression = qasm3_ast.FloatLiteral(init_value) - if self._check_only: - return [] + self._handle_extern_function_cleanup(statements, statement) return statements # pylint: disable=too-many-branches, too-many-statements, too-many-locals + @check_only_return_empty def _visit_classical_declaration( self, statement: qasm3_ast.ClassicalDeclaration ) -> list[qasm3_ast.Statement]: @@ -1983,7 +1968,6 @@ def _visit_classical_declaration( # populate the variable if statement.init_expression: - global_scope = self._scope_manager.get_global_scope() PulseValidator.validate_duration_or_stretch_statements( statement=statement, @@ -2113,16 +2097,13 @@ def _visit_classical_declaration( statement.init_expression = PulseValidator.make_complex_binary_expression(init_value) if isinstance(statement.init_expression, qasm3_ast.FunctionCall): - statement.init_expression = ( - self._handle_function_init_expression(statement.init_expression, init_value) - or statement.init_expression - ) - - if self._check_only: - return [] + function_name = statement.init_expression.name.name + if function_name in FUNCTION_MAP and isinstance(init_value, (float, int)): + statement.init_expression = qasm3_ast.FloatLiteral(init_value) return statements + @check_only_return_empty def _visit_classical_assignment( self, statement: qasm3_ast.ClassicalAssignment ) -> list[qasm3_ast.Statement]: @@ -2287,16 +2268,12 @@ def _visit_classical_assignment( ) if isinstance(statement.rvalue, qasm3_ast.FunctionCall): - statement.rvalue = ( - self._handle_function_init_expression(statement.rvalue, rvalue_eval) - or statement.rvalue - ) + function_name = statement.rvalue.name.name + if function_name in FUNCTION_MAP and isinstance(rvalue_eval, (float, int)): + statement.rvalue = qasm3_ast.FloatLiteral(rvalue_eval) self._handle_extern_function_cleanup(statements, statement) - if self._check_only: - return [] - return statements def _evaluate_array_initialization( @@ -2341,6 +2318,7 @@ def _update_branching_gate_depths(self) -> None: self._is_branch_clbits.clear() self._is_branch_qubits.clear() + @check_only_return_empty def _visit_branching_statement( self, statement: qasm3_ast.BranchingStatement ) -> list[qasm3_ast.Statement]: @@ -2470,9 +2448,6 @@ def ravel(bit_ind): if not self._in_branching_statement: self._update_branching_gate_depths() - if self._check_only: - return [] - return result # type: ignore[return-value] def _visit_forin_loop(self, statement: qasm3_ast.ForInLoop) -> list[qasm3_ast.Statement]: @@ -2559,6 +2534,7 @@ def _visit_forin_loop(self, statement: qasm3_ast.ForInLoop) -> list[qasm3_ast.St return [] return result + @check_only_return_empty def _visit_subroutine_definition( self, statement: qasm3_ast.SubroutineDefinition | qasm3_ast.ExternDeclaration ) -> Sequence[None | qasm3_ast.ExternDeclaration]: @@ -2608,8 +2584,6 @@ def _visit_subroutine_definition( statements.append(statement) self._subroutine_defns[fn_name] = statement - if self._check_only: - return [] return statements @@ -2823,7 +2797,7 @@ def _visit_alias_statement(self, statement: qasm3_ast.AliasStatement) -> list[No # this will only build a global alias map - # whenever we are referring to qubits , we will first check in the global map of registers + # whenever we are referring to qubits, we will first check in the global map of registers # if the register is present, we will use the global map to get the qubit labels # if not, we will check the alias map for the labels @@ -2926,8 +2900,7 @@ def _visit_switch_statement( # type: ignore[return] statement (SwitchStatement): The switch statement to visit. Returns: - list[Statement]: The list of statements generated by the switch statement, - or an empty list if self._check_only is True. + list[Statement]: The list of statements generated by the switch statement. """ # 1. analyze the target - it should ONLY be int, not casted switch_target = statement.target @@ -2972,8 +2945,6 @@ def _evaluate_case(statements): self._scope_manager.pop_scope() self._scope_manager.restore_context() - if self._check_only: - return [] return result case_fulfilled = False @@ -3135,6 +3106,7 @@ def _is_verbatim_pragma(statement: qasm3_ast.Pragma) -> bool: """ return statement.command.split() == ["braket", "verbatim"] + @check_only_return_empty def _visit_pragma(self, statement: qasm3_ast.Pragma) -> list[qasm3_ast.Pragma]: """ Visit a Pragma statement. @@ -3155,9 +3127,6 @@ def _visit_pragma(self, statement: qasm3_ast.Pragma) -> list[qasm3_ast.Pragma]: if self._is_verbatim_pragma(statement): self._verbatim_pragma_pending = True - if self._check_only: - return [] - return [statement] def _visit_box_statement(self, statement: qasm3_ast.Box) -> list[qasm3_ast.Statement]: @@ -3484,6 +3453,7 @@ def _visit_calibration_grammar_declaration( return [statement] + @check_only_return_empty def _visit_include(self, include: qasm3_ast.Include) -> list[qasm3_ast.Statement]: """Visit an include statement element. @@ -3500,8 +3470,6 @@ def _visit_include(self, include: qasm3_ast.Include) -> list[qasm3_ast.Statement f"File '{filename}' already included", error_node=include, span=include.span ) self._included_files.add(filename) - if self._check_only: - return [] return [include]