Skip to content
Merged
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
50 changes: 35 additions & 15 deletions src/sidewinder/analysis/transform/transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,7 +90,7 @@ def visit_Delete(self, node: ast.Delete) -> Any:
return node


def visit_Assert(self, node: ast.Assert) -> ast.Expr:
def visit_Assert(self, node: ast.Assert) -> list[ast.stmt]:
"""
Transform assert statement to explicit sidewinder call.

Expand All @@ -102,14 +102,20 @@ def visit_Assert(self, node: ast.Assert) -> ast.Expr:
becomes:
__sidewinder_assert__(expr<transformed>, msg<transformed>, __sidewinder_state=__sidewinder_state)
"""
args = [self._visit_expr(node.test)]
lowered_test = self._visit_expr(node.test)
args = [lowered_test.expr]
if node.msg is not None:
args.append(self._visit_expr(node.msg))
lowered_msg = self._visit_expr(node.msg)
args.append(lowered_msg.expr)

return ast.Expr(
ret: list[ast.stmt] = []
ret.extend(lowered_test.stmts)
ret.extend(lowered_msg.stmts)
ret.append(ast.Expr(
self._emit_hook_call(SidewinderHookNames.SIDEWINDER_ASSERT, *args),
lineno=0, col_offset=0,
)
))
return ret

def visit_Import(self, node: ast.Import) -> Any:
"""Import statements pass through unchanged."""
Expand All @@ -127,13 +133,16 @@ def visit_Nonlocal(self, node: ast.Nonlocal) -> Any:
"""Nonlocal statements pass through unchanged."""
return node

def visit_Expr(self, node: ast.Expr) -> ast.Expr:
def visit_Expr(self, node: ast.Expr) -> list[ast.stmt]:
"""
Transform expression statement — a bare expression used as a statement.
e.g. function calls whose return value is discarded: foo(), obj.method()
"""
node.value = self._visit_expr(node.value)
return node
lowered_value = self._visit_expr(node.value)
ret = []
ret.extend(lowered_value.stmts)
ret.append(ast.Expr(value=lowered_value.expr))
return ret

def visit_Pass(self, node: ast.Pass) -> Any:
"""Pass statements are unchanged."""
Expand Down Expand Up @@ -211,16 +220,27 @@ def visit_arguments(self, node: ast.arguments) -> Any:
# Arguments are handled by FunctionDef/Lambda visitors
return node

def visit_arg(self, node: ast.arg) -> Any:
def visit_arg(self, node: ast.arg) -> list[ast.AST]:
"""Transform function argument."""
# TODO: This is not a expr or stmt, if theres a complex expression here it may :blows raspberry:. (Pydantic annotations could complain)
if node.annotation:
node.annotation = self._visit_expr(node.annotation)
return node

def visit_keyword(self, node: ast.keyword) -> Any:
lowered_annotation = self._visit_expr(node.annotation)
node.annotation = lowered_annotation.expr
ret = []
ret.extend(lowered_annotation.stmts)
ret.append(node)
return ret
else:
return [node]

def visit_keyword(self, node: ast.keyword) -> list[ast.stmt]:
"""Transform keyword argument."""
node.value = self._visit_expr(node.value)
return node
lowered_keyword = self._visit_expr(node.value)
node.value = lowered_keyword.expr
ret = []
ret.extend(lowered_keyword.stmts)
ret.append(node)
return ret

def visit_alias(self, node: ast.alias) -> Any:
"""Import alias is unchanged."""
Expand Down
27 changes: 20 additions & 7 deletions src/sidewinder/analysis/transform/transformer_assign.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,21 +7,27 @@
class SidewinderAssignTransformerMixin(SidewinderTransformerHelpers):
def visit_Assign(self, node: ast.Assign) -> list[ast.stmt]:
"""Transform assignment statement - transform value."""
transformed_value = self._visit_expr(node.value)
lowered_value = self._visit_expr(node.value)
if len(node.targets) != 1:
raise NotImplementedError(f"Sidewinder currently only supports a single target for assign statements like {ast.unparse(node)}")
to_return = []
to_return.extend(lowered_value.stmts)
for target in node.targets:
to_return.extend(self._visit_target(target, transformed_value))
to_return.extend(self._visit_target(target, lowered_value.expr))
return to_return

def visit_AnnAssign(self, node: ast.AnnAssign) -> list[ast.stmt]:
"""Transform annotated assignment statement."""
# Transform the annotation (it's an expr)
node.annotation = self._visit_expr(node.annotation)
stmts = []
lowered_annotation = self._visit_expr(node.annotation)
node.annotation = lowered_annotation.expr
stmts.extend(lowered_annotation.stmts)
# Transform the value if it exists (can be None)
if node.value is not None:
return self._visit_target(node.target, self._visit_expr(node.value))
lowered_value = self._visit_expr(node.value)
stmts.append(lowered_value.stmts)
return stmts + self._visit_target(node.target, lowered_value.expr)
return []

def visit_AugAssign(self, node: ast.AugAssign) -> list[ast.stmt]:
Expand All @@ -47,15 +53,22 @@ def visit_AugAssign(self, node: ast.AugAssign) -> list[ast.stmt]:

method_name = f'__sidewinder_i{op_name}__'

to_return = []

read_target = copy.deepcopy(node.target)
lowered_read_target = self._visit_expr(read_target)
to_return.extend(lowered_read_target.stmts)
read_target.ctx = ast.Load()

lowered_value = self._visit_expr(node.value)
to_return.extend(lowered_value.stmts)
rhs_value = ast.Call(
func=ast.Attribute(
value=self._visit_expr(read_target),
value=lowered_read_target.expr,
attr=method_name,
ctx=ast.Load(),
),
args=[self._visit_expr(node.value)],
args=[lowered_value.expr],
keywords=[
ast.keyword(
arg='__sidewinder_state',
Expand All @@ -65,7 +78,7 @@ def visit_AugAssign(self, node: ast.AugAssign) -> list[ast.stmt]:
)

tmp = self._fresh_temp()
to_return = []

to_return.append(
ast.Assign(
targets=[ast.Name(id=tmp, ctx=ast.Store())],
Expand Down
7 changes: 7 additions & 0 deletions src/sidewinder/analysis/transform/transformer_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,3 +34,10 @@ def visit(self, node: ast.AST) -> ast.AST: ...
def visit(self, node: ast.AST) -> Any:
return super().visit(node)

class LoweredExpr:
stmts: list[ast.stmt]
expr: ast.expr

def __init__(self, stmts, expr):
self.stmts = stmts
self.expr = expr
45 changes: 35 additions & 10 deletions src/sidewinder/analysis/transform/transformer_classes.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,21 +4,46 @@
from sidewinder.analysis.transform.transformer_helpers import SidewinderTransformerHelpers

class SidewinderClassTransformerMixin(SidewinderTransformerHelpers):
def visit_ClassDef(self, node: ast.ClassDef) -> Any:
def visit_ClassDef(self, node: ast.ClassDef) -> list[ast.stmt]:
"""Visit a class definition - transform all methods."""
self.current_scope.append(node.name)

# Transform class body
with self.current_context.enter_context(node, "body") as c:
with self.current_context.enter_context(node, "body"):
node.body = self._visit_list_of_stmts(node.body)

# Transform decorators
node.decorator_list = [self._visit_expr(dec) for dec in node.decorator_list]

# Transform bases
node.bases = [self._visit_expr(bas) for bas in node.bases]
new_bases = []
new_bases_stmts: list[ast.stmt] = []
for base in node.bases:
lowered = self._visit_expr(base)
# TODO: handle base expression stmts if needed
new_bases.append(lowered.expr)
new_bases_stmts.extend(lowered.stmts)

node.bases = new_bases

# Transform keywords
for kw in node.keywords:
kw.value = self._visit_expr(kw.value)

lowered = self._visit_expr(kw.value)
# TODO: handle keyword expression stmts if needed
kw.value = lowered.expr

self.current_scope.pop()
return node

# Lower decorators
decorator_pre, decorator_post = self._transform_decorators(node)

ret: list[ast.stmt] = []
ret.extend(new_bases_stmts)

ret.extend(decorator_pre)

# Emit class definition
ret.append(node)

# Apply decorators after class exists
decorator_transforms = self._visit_list_of_stmts(decorator_post)
ret.extend(decorator_transforms)

return ret
19 changes: 7 additions & 12 deletions src/sidewinder/analysis/transform/transformer_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ class TransformerContext:
def __init__(self):
# Stack of (node, attr_name) tuples
# e.g., (If node, 'body') or (If node, 'orelse')
self.context_stack: List[Tuple[ast.AST, str]] = []
self.context_stack: List[Tuple[ast.AST, str, list]] = []

def push_context(self, node: ast.AST, attr_name: str) -> None:
"""
Expand All @@ -28,13 +28,14 @@ def push_context(self, node: ast.AST, attr_name: str) -> None:
node: The AST node containing the statement list
attr_name: The attribute name (e.g., 'body', 'orelse', 'finalbody')
"""
self.context_stack.append((node, attr_name))
self.context_stack.append((node, attr_name, []))

def pop_context(self) -> Tuple[ast.AST, str]:
"""Exit the current context."""
if not self.context_stack:
raise RuntimeError("Cannot pop_context: context stack is empty")
return self.context_stack.pop()
# setattr(self.context_stack[-1][0], self.context_stack[-1][1], self.context_stack[-1][2])
return self.context_stack.pop()[:-1]

@contextmanager
def enter_context(self, node: ast.AST, attr_name: str):
Expand All @@ -54,10 +55,6 @@ def enter_context(self, node: ast.AST, attr_name: str):
finally:
self.pop_context()

def current_context(self) -> Optional[Tuple[ast.AST, str]]:
"""Get the current (node, attr_name) context."""
return self.context_stack[-1] if self.context_stack else None

def append_stmt(self, stmt: ast.stmt) -> None:
"""
Append a statement to the current context's statement list.
Expand All @@ -67,16 +64,14 @@ def append_stmt(self, stmt: ast.stmt) -> None:
if not self.context_stack:
raise RuntimeError("No context available - cannot append statement")

node, attr_name = self.context_stack[-1]
node, attr_name, stmt_list = self.context_stack[-1]

if not hasattr(node, attr_name):
raise RuntimeError(
f"Node {type(node).__name__} has no attribute '{attr_name}'"
)

stmt_list = getattr(node, attr_name)

if not isinstance(stmt_list, list):

if not isinstance(getattr(node, attr_name), list):
raise RuntimeError(
f"Attribute {attr_name} of {type(node).__name__} is not a list"
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,14 @@
class SidewinderControlFlowBreakerTransformerMixin(SidewinderTransformerHelpers):
def visit_Return(self, node: ast.Return) -> Any:
"""Transform return statement - transform the return value."""
return ast.Expr(value=self._emit_hook_call(
lowered_value = self._visit_expr(node.value) if node.value else None
ret = []
ret.extend(lowered_value.stmts) if lowered_value else None
ret.append(ast.Expr(value=self._emit_hook_call(
SidewinderHookNames.SIDEWINDER_RETURN,
self._visit_expr(node.value) if node.value else ast.Constant(value=None)
), lineno=0, col_offset=0)
lowered_value.expr if lowered_value else ast.Constant(value=None)
), lineno=0, col_offset=0))
return ret

def visit_Break(self, node: ast.Break) -> ast.Expr:
"""
Expand Down Expand Up @@ -55,11 +59,16 @@ def visit_Raise(self, node: ast.Raise) -> ast.Expr:
__sidewinder_raise__(__sidewinder_state=__sidewinder_state)
"""

stmts = []
args = []
if node.exc is not None:
args = [self._visit_expr(node.exc)]
lowered_exc = self._visit_expr(node.exc)
args = [lowered_exc.expr]
stmts.extend(lowered_exc.stmts)
if node.cause is not None:
args.append(self._visit_expr(node.cause))
lowered_cause = self._visit_expr(node.cause)
args.append(lowered_cause.expr)
stmts.extend(lowered_cause.stmts)

return ast.Expr(
value=self._emit_hook_call(SidewinderHookNames.SIDEWINDER_RAISE, *args),
Expand Down
Loading
Loading