diff --git a/CHANGELOG.md b/CHANGELOG.md index 62492ba..46ef72f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,7 @@ Types of changes: ### Added ### 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 888f818..be8a90b 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 d837f14..5d26614 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: @@ -205,6 +222,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]: @@ -286,8 +304,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 @@ -302,8 +318,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 = [] @@ -565,26 +580,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: @@ -753,7 +748,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: @@ -797,12 +791,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) ): @@ -842,7 +837,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( @@ -1172,6 +1166,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, @@ -1270,9 +1265,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: @@ -1287,6 +1279,7 @@ def _visit_continue(self, statement: qasm3_ast.ContinueStatement) -> None: error_node=statement, ) + @check_only_return_empty def _visit_custom_gate_operation( self, operation: qasm3_ast.QuantumGate, @@ -1405,11 +1398,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, @@ -1498,11 +1489,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, @@ -1562,9 +1552,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 @@ -1734,9 +1721,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]: @@ -1847,18 +1834,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]: @@ -1958,7 +1943,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, @@ -2088,16 +2072,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]: @@ -2262,16 +2243,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( @@ -2316,6 +2293,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]: @@ -2445,9 +2423,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]: @@ -2534,6 +2509,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]: @@ -2583,8 +2559,6 @@ def _visit_subroutine_definition( statements.append(statement) self._subroutine_defns[fn_name] = statement - if self._check_only: - return [] return statements @@ -2798,7 +2772,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 @@ -2901,8 +2875,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 @@ -2947,8 +2920,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 @@ -3110,6 +3081,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. @@ -3130,9 +3102,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]: @@ -3459,6 +3428,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. @@ -3475,8 +3445,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]