diff --git a/src/pyqasm/elements.py b/src/pyqasm/elements.py index b75c740..a1336a6 100644 --- a/src/pyqasm/elements.py +++ b/src/pyqasm/elements.py @@ -47,6 +47,13 @@ def is_internal_qubit_register(qubit_name: str) -> bool: ) +INTERNAL_QUANTUM_ARGUMENT = "__(QUANTUM_ARGUMENT)__" +"""Reserved string used to prevent shadowing for parameters in function definitions. +The register appears as a suffix to a variable name ("variable__(QUANTUM_ARGUMENT)"). +A user is not able to define a function parameter with this suffix due to the parentheses. +""" + + class InversionOp(Enum): """ Enum for specifying the inversion action of a gate. diff --git a/src/pyqasm/subroutines.py b/src/pyqasm/subroutines.py index 8e8f197..9ccf990 100644 --- a/src/pyqasm/subroutines.py +++ b/src/pyqasm/subroutines.py @@ -45,7 +45,7 @@ from openqasm3.printer import dumps from pyqasm.analyzer import Qasm3Analyzer -from pyqasm.elements import Variable +from pyqasm.elements import INTERNAL_QUANTUM_ARGUMENT, Variable from pyqasm.exceptions import ValidationError, raise_qasm3_error from pyqasm.expressions import Qasm3ExprEvaluator from pyqasm.transformer import Qasm3Transformer @@ -521,6 +521,12 @@ def process_quantum_arg( # pylint: disable=too-many-locals """ actual_arg_name = Qasm3SubroutineProcessor.get_fn_actual_arg_name(actual_arg) formal_reg_name = formal_arg.name.name + internal_reg_name = formal_reg_name + # If our actual variable is the same as the function argument, + # give the function argument a temporary name for internal use + if actual_arg_name == formal_reg_name: + internal_reg_name += INTERNAL_QUANTUM_ARGUMENT + formal_qubit_size = Qasm3ExprEvaluator.evaluate_expression( formal_arg.size, reqd_type=IntType, const_expr=True )[0] @@ -536,7 +542,7 @@ def process_quantum_arg( # pylint: disable=too-many-locals error_node=fn_defn.arguments, span=formal_arg.span, ) - formal_qreg_size_map[formal_reg_name] = formal_qubit_size + formal_qreg_size_map[internal_reg_name] = formal_qubit_size # we expect that actual arg is qubit type only # note that we ONLY check in global scope as @@ -603,10 +609,10 @@ def process_quantum_arg( # pylint: disable=too-many-locals ) for idx, qid in enumerate(resolved_qids): - qubit_transform_map[(formal_reg_name, idx)] = (resolved_reg_name, qid) + qubit_transform_map[(internal_reg_name, idx)] = (resolved_reg_name, qid) return Variable( - name=formal_reg_name, + name=internal_reg_name, base_type=QubitDeclaration, base_size=formal_qubit_size, dims=None, diff --git a/src/pyqasm/visitor.py b/src/pyqasm/visitor.py index fd4d0dd..ef9fda1 100644 --- a/src/pyqasm/visitor.py +++ b/src/pyqasm/visitor.py @@ -34,6 +34,7 @@ from pyqasm.analyzer import Qasm3Analyzer from pyqasm.elements import ( + INTERNAL_QUANTUM_ARGUMENT, INTERNAL_QUBIT_REGISTER, Capture, ClbitDepthNode, @@ -1499,11 +1500,24 @@ def _visit_generic_gate_operation( # pylint: disable=too-many-branches, too-man for transform_map, size_map in zip( reversed(self._function_qreg_transform_map), reversed(self._function_qreg_size_map) ): - operation.qubits = ( - Qasm3Transformer.transform_function_qubits( # type: ignore [assignment] - operation, transform_map, size_map + try: + operation.qubits = ( + Qasm3Transformer.transform_function_qubits( # type: ignore [assignment] + operation, transform_map, size_map + ) + ) + except KeyError: + for qubit in operation.qubits: + # Each qubit may be an IndexedIdentifier or an Identifier + if isinstance(qubit, qasm3_ast.IndexedIdentifier): + qubit.name.name += INTERNAL_QUANTUM_ARGUMENT + else: + qubit.name += INTERNAL_QUANTUM_ARGUMENT + operation.qubits = ( + Qasm3Transformer.transform_function_qubits( # type: ignore [assignment] + operation, transform_map, size_map + ) ) - ) operation.qubits = self._get_op_bits(operation, qubits=True) # type: ignore