Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 4 additions & 5 deletions src/pyqasm/pulse/validator.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,11 +114,10 @@ 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.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

G1 — new line-too-long, fails CI.

src/pyqasm/pulse/validator.py:118:0: C0301: Line too long (105/100)

Adding (Any) and the trailing period (N5/N6) pushed it over. Wrapping the clause onto a continuation line matches how the rest of the file handles long Args: entries:

base_type (Any): The declared type; the function does nothing if it is
    neither DurationType nor StretchType.

This is the only remaining CI failure — visitor.py is clean, and mypy/black/isort all pass.

rvalue (Any): The initializer or assigned value.
global_scope (dict): Global symbol table.

Raises:
ValidationError: If the assigned value is not a DurationLiteral,
Expand Down
114 changes: 43 additions & 71 deletions src/pyqasm/visitor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -83,6 +84,22 @@
logger = logging.getLogger(__name__)
logger.propagate = False

F = TypeVar("F", bound=Callable[..., Any])


def check_only_return_empty(func: F) -> F:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Confirming N1/N2 as fixed. F -> F correctly preserves the wrapped signature (the earlier Callable[..., Any] would have erased it), and # type: ignore on the return is the standard escape for this pattern — mypy is clean.

Worth noting for whoever extends this: the decorator is only safe on methods whose every return should be suppressed under check_only. That's what R1 turned on. If a fourth method gets decorated later, the check is "does it have an early return that isn't itself gated" — a one-line AST sweep catches it.

"""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:
Expand Down Expand Up @@ -199,6 +216,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]:
Expand Down Expand Up @@ -280,8 +298,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
Expand Down Expand Up @@ -530,26 +546,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:
Expand Down Expand Up @@ -718,7 +714,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:
Comment thread
micpap25 marked this conversation as resolved.
Expand Down Expand Up @@ -762,12 +757,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)
):
Expand Down Expand Up @@ -807,7 +803,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(
Expand Down Expand Up @@ -1137,6 +1132,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,
Expand Down Expand Up @@ -1235,9 +1231,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:
Expand All @@ -1252,6 +1245,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,
Expand Down Expand Up @@ -1352,11 +1346,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,
Expand Down Expand Up @@ -1428,11 +1420,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,
Expand Down Expand Up @@ -1492,9 +1483,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
Expand Down Expand Up @@ -1659,9 +1647,9 @@ def _visit_generic_gate_operation( # pylint: disable=too-many-branches, too-man

if self._check_only:
return []

return result

@check_only_return_empty
def _visit_constant_declaration(
self, statement: qasm3_ast.ConstantDeclaration
) -> list[qasm3_ast.Statement]:
Expand Down Expand Up @@ -1772,18 +1760,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)):
Comment thread
ryanhill1 marked this conversation as resolved.
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]:
Expand Down Expand Up @@ -1883,7 +1869,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,
Expand Down Expand Up @@ -2013,16 +1998,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]:
Expand Down Expand Up @@ -2187,16 +2169,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(
Expand Down Expand Up @@ -2241,6 +2219,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]:
Expand Down Expand Up @@ -2370,9 +2349,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]:
Expand Down Expand Up @@ -2459,6 +2435,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]:
Expand Down Expand Up @@ -2508,8 +2485,6 @@ def _visit_subroutine_definition(
statements.append(statement)

self._subroutine_defns[fn_name] = statement
if self._check_only:
return []

return statements

Expand Down Expand Up @@ -2724,7 +2699,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
Expand Down Expand Up @@ -2873,8 +2848,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
Expand Down Expand Up @@ -3384,6 +3357,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.

Expand All @@ -3400,8 +3374,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]

Expand Down
Loading